Spec-Zone.ru › PyTorch 2.14

KLDivLoss

class torch.nn.KLDivLoss(size_average=None, reduce=None, reduction='mean', log_target=False) [исходный код]

Функция потерь на основе дивергенции Кульбака — Лейблера.

Для тензоров одинаковой формы ypred,ytruey_{\text{pred}},\ y_{\text{true}}, где ypredy_{\text{pred}} — это input, а ytruey_{\text{true}} — это target, определим поэлементную дивергенцию Кульбака — Лейблера следующим образом:

L(ypred,ytrue)=ytrue⋅log⁡ytrueypred=ytrue⋅(log⁡ytrue−log⁡ypred)L(y_{\text{pred}},\ y_{\text{true}}) = y_{\text{true}} \cdot \log \frac{y_{\text{true}}}{y_{\text{pred}}} = y_{\text{true}} \cdot (\log y_{\text{true}} - \log y_{\text{pred}})

Чтобы избежать проблем с потерей точности при вычислении этой величины, функция потерь ожидает, что аргумент input будет задан в логарифмическом пространстве. Аргумент target также можно задать в логарифмическом пространстве, если log_target= True.

Иными словами, эта функция примерно эквивалентна вычислению

if not log_target:  # default
    loss_pointwise = target * (target.log() - input)
else:
    loss_pointwise = target.exp() * (target - input)

с последующим сведением результата в зависимости от аргумента reduction следующим образом:

if reduction == "mean":  # default
    loss = loss_pointwise.mean()
elif reduction == "batchmean":  # mathematically correct
    loss = loss_pointwise.sum() / input.size(0)
elif reduction == "sum":
    loss = loss_pointwise.sum()
else:  # reduction == "none"
    loss = loss_pointwise

Примечание

Как и все остальные функции потерь в PyTorch, эта функция ожидает, что первый аргумент, input, будет выходными данными модели (например, нейронной сети), а второй, target, — наблюдениями из набора данных. Это отличается от стандартной математической записи KL(P∣∣Q)KL(P\ ||\ Q), где PP обозначает распределение наблюдений, а QQ — модель.

Предупреждение

reduction= “mean” не возвращает истинное значение дивергенции Кульбака — Лейблера; используйте reduction= “batchmean”, соответствующий математическому определению.

Параметры:
  • size_average (bool, необязательно) – Устаревший параметр (см. reduction). По умолчанию потери усредняются по каждому элементу потерь в пакете. Обратите внимание, что для некоторых функций потерь на один пример приходится несколько элементов. Если поле size_average задано как False, вместо этого потери суммируются для каждого мини-пакета. Игнорируется, если reduce равно False. По умолчанию: True
  • reduce (bool, необязательно) – Устаревший параметр (см. reduction). По умолчанию потери усредняются или суммируются по наблюдениям для каждого мини-пакета в зависимости от size_average. Если reduce равно False, возвращает потери для каждого элемента пакета и игнорирует size_average. По умолчанию: True
  • reduction (str, необязательно) – Задает способ сведения выходных данных. По умолчанию: “mean”
  • log_target (bool, необязательно) – Указывает, задан ли target в логарифмическом пространстве. По умолчанию: False
Форма:
  • Входные данные: (∗)(*), где ∗* означает любое количество измерений.
  • Целевые данные: (∗)(*), той же формы, что и входные данные.
  • Выходные данные: по умолчанию скаляр. Если reduction равно ‘none’, то (∗)(*), той же формы, что и входные данные.

Примеры

>>> kl_loss = nn.KLDivLoss(reduction="batchmean")
>>> # input should be a distribution in the log space
>>> input = F.log_softmax(torch.randn(3, 5, requires_grad=True), dim=1)
>>> # Sample a batch of distributions. Usually this would come from the dataset
>>> target = F.softmax(torch.rand(3, 5), dim=1)
>>> output = kl_loss(input, target)
>>>
>>> kl_loss = nn.KLDivLoss(reduction="batchmean", log_target=True)
>>> log_target = F.log_softmax(torch.rand(3, 5), dim=1)
>>> output = kl_loss(input, log_target)
forward(input, target) [исходный код]

Выполняет прямой проход.

Тип возвращаемого значения:

Tensor

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.KLDivLoss.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API