Класс CrossEntropyLoss
Класс CrossEntropyLoss вычисляет потерю для задач многоклассовой классификации.
Он комбинирует функцию LogSoftmax и класс NLLLoss в одном слое,
что делает его более стабильным численно. Принимает на вход логиты модели и целевые классы.
Первым параметром можно указать вес классов, вторым - размер батча для усреднения.
Синтаксис
torch.nn.CrossEntropyLoss(weight=None, reduction='mean')
Пример
Создадим простую модель и вычислим потерю для трех классов:
import torch
import torch.nn as nn
# логиты для двух образцов и трех классов
logits = torch.tensor([
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0]
])
targets = torch.tensor([2, 1]) # истинные классы
loss_fn = nn.CrossEntropyLoss()
loss = loss_fn(logits, targets)
print(loss)
Результат выполнения кода:
tensor(0.9519)
Пример
Используем параметр reduction для получения суммы потерь:
import torch
import torch.nn as nn
torch.manual_seed(0)
logits = torch.randn(4, 5) # 4 образца, 5 классов
targets = torch.randint(0, 5, (4,))
loss_fn = nn.CrossEntropyLoss(reduction='sum')
loss = loss_fn(logits, targets)
print(loss)
Результат выполнения кода:
tensor(9.4219)
Пример
Передаем веса для классов, чтобы сбалансировать данные:
import torch
import torch.nn as nn
torch.manual_seed(0)
logits = torch.randn(3, 4)
targets = torch.tensor([0, 1, 3])
# вес для каждого класса
class_weight = torch.tensor([0.5, 1.0, 2.0, 0.8])
loss_fn = nn.CrossEntropyLoss(weight=class_weight)
loss = loss_fn(logits, targets)
print(loss)
Результат выполнения кода:
tensor(2.1852)
Пример
Применяем потерю в процессе обучения модели:
import torch
import torch.nn as nn
import torch.optim as optim
torch.manual_seed(0)
model = nn.Linear(10, 3)
loss_fn = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.01)
x = torch.randn(8, 10)
y = torch.randint(0, 3, (8,))
for epoch in range(5):
optimizer.zero_grad()
output = model(x)
loss = loss_fn(output, y)
loss.backward()
optimizer.step()
print(f"Epoch {epoch+1}: loss = {loss.item():.4f}")
Результат выполнения кода:
Epoch 1: loss = 1.2807
Epoch 2: loss = 1.1963
Epoch 3: loss = 1.1304
Epoch 4: loss = 1.0790
Epoch 5: loss = 1.0388
Смотрите также
-
класс
NLLLoss,
который вычисляет отрицательную логарифмическую вероятность -
класс
BCEWithLogitsLoss,
который использует сигмоиду и бинарную кросс-энтропию -
класс
MSELoss,
который вычисляет среднеквадратичную ошибку для регрессии -
класс
KLDivLoss,
который вычисляет дивергенцию Кульбака-Лейблера