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

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