Функция F.binary_cross_entropy_with_logits
Функция F.binary_cross_entropy_with_logits вычисляет бинарную кросс-энтропию между предсказаниями модели и целевыми метками для задач бинарной классификации. Она принимает логиты (необработанные выходы модели) и применяет к ним сигмоиду внутри себя, что обеспечивает численную стабильность и эффективность. Первым параметром функция принимает тензор логитов, вторым - тензор целевых меток. Третьим параметром можно передать веса для каждого элемента выборки.
Синтаксис
torch.nn.functional.binary_cross_entropy_with_logits(input, target, weight=None, reduction='mean', pos_weight=None)
В качестве параметров функция принимает:
-
input- тензор логитов произвольной формы -
target- тензор целевых меток той же формы, что иinput -
weight- веса для каждого элемента выборки -
reduction- способ агрегации потерь: 'none', 'mean', 'sum' -
pos_weight- веса для положительного класса, используется для балансировки классов
Пример
Давайте вычислим бинарную кросс-энтропию для простых логитов и целевых меток:
import torch
import torch.nn.functional as F
logits = torch.tensor([2.0, -1.5, 0.0, 3.2])
target = torch.tensor([1.0, 0.0, 1.0, 1.0])
loss = F.binary_cross_entropy_with_logits(logits, target)
print(loss)
Результат выполнения кода:
tensor(0.3231)
Пример
Используем параметр weight для взвешивания потерь:
import torch
import torch.nn.functional as F
logits = torch.tensor([0.8, -2.1, 1.5, -0.3])
target = torch.tensor([1.0, 0.0, 1.0, 0.0])
weights = torch.tensor([0.5, 1.0, 1.5, 0.8])
loss = F.binary_cross_entropy_with_logits(logits, target, weight=weights)
print(loss)
Результат выполнения кода:
tensor(0.1749)
Пример
Параметр pos_weight позволяет балансировать положительный и отрицательный классы:
import torch
import torch.nn.functional as F
logits = torch.tensor([[0.5], [-1.0], [2.0], [-0.5]])
target = torch.tensor([[1.0], [0.0], [1.0], [0.0]])
pos_weight = torch.tensor([2.0])
loss = F.binary_cross_entropy_with_logits(logits, target, pos_weight=pos_weight)
print(loss)
Результат выполнения кода:
tensor(0.2252)
Пример
Выполним вычисления без агрегации с параметром reduction='none':
import torch
import torch.nn.functional as F
logits = torch.tensor([[1.2, -0.8], [0.3, 2.1]])
target = torch.tensor([[0.0, 1.0], [1.0, 0.0]])
loss = F.binary_cross_entropy_with_logits(logits, target, reduction='none')
print(loss)
Результат выполнения кода:
tensor([
[1.1269, 0.3711],
[0.5544, 2.1426],
])
Пример
Используем функцию в контексте обучения модели с фиксированным зерном для воспроизводимости:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
batch_size = 3
num_features = 2
logits = torch.randn(batch_size, num_features)
target = torch.randint(0, 2, (batch_size, num_features)).float()
loss = F.binary_cross_entropy_with_logits(logits, target)
print(loss)
Результат выполнения кода:
tensor(1.2762)
Смотрите также
-
функцию
binary_cross_entropy,
которая вычисляет бинарную кросс-энтропию для вероятностей -
функцию
sigmoid,
которая преобразует логиты в вероятности -
функцию
cross_entropy,
которая вычисляет кросс-энтропию для многоклассовой классификации -
функцию
mse_loss,
которая вычисляет среднеквадратичную ошибку для регрессионных задач