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

Класс 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,
    который вычисляет дивергенцию Кульбака-Лейблера
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить