РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
432 of 769 menu

Функция F.nll_loss

Функция F.nll_loss вычисляет отрицательную логарифмическую вероятность (negative log likelihood loss). Она часто используется в задачах классификации вместе с функцией F.log_softmax. Первым параметром функция принимает тензор с логарифмическими вероятностями, вторым параметром - тензор с индексами классов. Дополнительно можно указать веса классов, параметр reduction и индекс игнорирования.

Синтаксис

torch.nn.functional.nll_loss(input, target, weight=None, reduction='mean', ignore_index=-100)

Пример

Давайте вычислим значение функции на простом примере с тремя классами и двумя образцами:

import torch import torch.nn.functional as F torch.manual_seed(0) log_probs = torch.log_softmax(torch.randn(2, 3), dim=1) target = torch.tensor([0, 2]) res = F.nll_loss(log_probs, target) print(res)

Результат выполнения кода:

tensor(0.8521)

Пример

Теперь рассмотрим пример с указанием весов классов и параметром reduction равным 'sum':

import torch import torch.nn.functional as F torch.manual_seed(0) log_probs = torch.log_softmax(torch.randn(2, 3), dim=1) target = torch.tensor([0, 2]) weights = torch.tensor([0.5, 1.0, 2.0]) res = F.nll_loss(log_probs, target, weight=weights, reduction='sum') print(res)

Результат выполнения кода:

tensor(1.7042)

Пример

Рассмотрим использование параметра ignore_index для исключения определённого класса из расчёта:

import torch import torch.nn.functional as F torch.manual_seed(0) log_probs = torch.log_softmax(torch.randn(2, 3), dim=1) target = torch.tensor([0, -100]) res = F.nll_loss(log_probs, target, ignore_index=-100) print(res)

Результат выполнения кода:

tensor(0.9458)

Смотрите также

  • функцию log_softmax,
    которая вычисляет логарифм от функции softmax
  • функцию cross_entropy,
    которая является комбинацией log_softmax и nll_loss
  • функцию mse_loss,
    которая вычисляет среднеквадратическую ошибку
  • функцию kl_div,
    которая вычисляет дивергенцию Кульбака-Лейблера
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить