Функция 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,
которая вычисляет бинарную кросс-энтропию