Функция binary_cross_entropy
Функция binary_cross_entropy из модуля torch.nn.functional вычисляет бинарную кросс-энтропию между предсказанными вероятностями и истинными метками. Первым параметром функция принимает тензор предсказаний input, вторым - тензор целевых значений target. Третьим параметром можно передать весовой коэффициент weight, а четвёртым - указание оси для редукции reduction. Функция возвращает скалярное значение или тензор потерь в зависимости от параметра reduction.
Синтаксис
torch.nn.functional.binary_cross_entropy(input, target, weight=None, reduction='mean')
Параметр input - тензор предсказанных вероятностей (значения должны быть в диапазоне от 0 до 1). Параметр target - тензор истинных меток (значения 0 или 1). Параметр weight - необязательный тензор весов для каждого элемента. Параметр reduction принимает значения 'none', 'mean', 'sum' и определяет способ агрегации потерь.
Пример
Давайте вычислим бинарную кросс-энтропию для простых предсказаний:
import torch
import torch.nn.functional as F
predictions = torch.tensor([0.9, 0.2, 0.8, 0.4])
targets = torch.tensor([1.0, 0.0, 1.0, 0.0])
loss = F.binary_cross_entropy(predictions, targets)
print(loss)
Результат выполнения кода:
tensor(0.2328)
Пример
Используем функцию с параметром reduction='none', чтобы получить потери для каждого элемента отдельно:
import torch
import torch.nn.functional as F
predictions = torch.tensor([0.9, 0.2, 0.8, 0.4])
targets = torch.tensor([1.0, 0.0, 1.0, 0.0])
loss = F.binary_cross_entropy(predictions, targets, reduction='none')
print(loss)
Результат выполнения кода:
tensor([0.1054, 0.2231, 0.2231, 0.5108])
Пример
Применим весовой коэффициент для разных классов:
import torch
import torch.nn.functional as F
predictions = torch.tensor([0.9, 0.2, 0.8, 0.4])
targets = torch.tensor([1.0, 0.0, 1.0, 0.0])
weights = torch.tensor([0.5, 1.5, 0.5, 1.5])
loss = F.binary_cross_entropy(predictions, targets, weight=weights)
print(loss)
Результат выполнения кода:
tensor(0.3628)
Пример
Рассмотрим использование в контексте обучения нейронной сети:
import torch
import torch.nn as nn
import torch.nn.functional as F
torch.manual_seed(0)
model = nn.Linear(5, 1)
x = torch.randn(4, 5)
targets = torch.tensor([[1.0], [0.0], [1.0], [0.0]])
logits = model(x)
predictions = torch.sigmoid(logits)
loss = F.binary_cross_entropy(predictions, targets)
print(loss.item())
Результат выполнения кода:
0.801130473613739
Смотрите также
-
функцию
binary_cross_entropy_with_logits,
которая объединяет сигмоиду и бинарную кросс-энтропию в одной функции -
функцию
cross_entropy,
которая вычисляет кросс-энтропию для многоклассовой классификации -
функцию
mse_loss,
которая вычисляет среднеквадратичную ошибку для регрессионных задач -
функцию
sigmoid,
которая преобразует логиты в вероятности для бинарной классификации