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

Функция cross_entropy

Функция cross_entropy вычисляет кросс-энтропийную функцию потерь между предсказанными вероятностями и истинными метками. Первым параметром функция принимает логиты модели (не нормализованные предсказания), вторым - истинные метки классов. Функция внутри себя применяет log_softmax и затем nll_loss, что делает её оптимальным выбором для многоклассовой классификации.

Синтаксис

torch.nn.functional.cross_entropy(input, target, weight=None, size_average=None, ignore_index=-100, reduce=None, reduction='mean', label_smoothing=0.0)

Пример

Давайте вычислим кросс-энтропию для трёх классов с одним объектом:

import torch import torch.nn.functional as F torch.manual_seed(0) logits = torch.tensor([[2.0, 1.0, 0.5]]) target = torch.tensor([0]) loss = F.cross_entropy(logits, target) print(loss)

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

tensor(0.4528)

Пример

Давайте вычислим кросс-энтропию для батча из трёх объектов с четырьмя классами:

import torch import torch.nn.functional as F torch.manual_seed(0) logits = torch.randn(3, 4) target = torch.tensor([1, 3, 0]) loss = F.cross_entropy(logits, target) print(loss)

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

tensor(1.2511)

Пример

Используем параметр reduction для получения суммы потерь вместо среднего значения:

import torch import torch.nn.functional as F torch.manual_seed(0) logits = torch.randn(2, 5) target = torch.tensor([2, 4]) loss = F.cross_entropy(logits, target, reduction='sum') print(loss)

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

tensor(3.6921)

Пример

Применим сглаживание меток для улучшения обобщения модели:

import torch import torch.nn.functional as F torch.manual_seed(0) logits = torch.tensor([[3.0, 1.0, 0.5]]) target = torch.tensor([0]) loss = F.cross_entropy(logits, target, label_smoothing=0.1) print(loss)

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

tensor(0.5818)

Пример

Используем веса классов для учёта дисбаланса в выборке:

import torch import torch.nn.functional as F torch.manual_seed(0) logits = torch.randn(4, 3) target = torch.tensor([0, 1, 2, 0]) class_weights = torch.tensor([0.5, 2.0, 1.5]) loss = F.cross_entropy(logits, target, weight=class_weights) print(loss)

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

tensor(1.5329)

Пример

Игнорируем определённый индекс меток при вычислении потерь:

import torch import torch.nn.functional as F torch.manual_seed(0) logits = torch.randn(3, 4) target = torch.tensor([1, -100, 0]) loss = F.cross_entropy(logits, target, ignore_index=-100) print(loss)

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

tensor(1.6172)

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

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