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

Класс BCELoss

Класс BCELoss вычисляет бинарную кросс-энтропию между целевыми значениями и предсказаниями модели. Входные данные должны быть вероятностями в диапазоне от 0 до 1, что обычно достигается применением сигмоиды к выходу модели. Потери вычисляются по формуле: loss = -w * (y * log(x) + (1 - y) * log(1 - x)).

Основные параметры класса:

  • weight - тензор весов для каждого элемента пакета. Позволяет задать важность отдельных примеров.
  • reduction - способ агрегации потерь. Может принимать значения 'none' (без агрегации), 'mean' (усреднение) или 'sum' (суммирование).
  • pos_weight - вес положительного класса для борьбы с дисбалансом классов.

Синтаксис

torch.nn.BCELoss(weight=None, reduction='mean', pos_weight=None)

Пример базового использования

Давайте создадим экземпляр класса и вычислим потери для простого случая бинарной классификации:

import torch import torch.nn as nn loss_fn = nn.BCELoss() predictions = torch.tensor([0.9, 0.2, 0.7, 0.4]) targets = torch.tensor([1.0, 0.0, 1.0, 0.0]) loss = loss_fn(predictions, targets) print(loss)

Результат выполнения кода:

tensor(0.2886)

Пример с использованием весов

Давайте зададим разные веса для каждого примера в пакете:

import torch import torch.nn as nn weights = torch.tensor([2.0, 1.0, 3.0, 1.0]) loss_fn = nn.BCELoss(weight=weights) predictions = torch.tensor([0.9, 0.2, 0.7, 0.4]) targets = torch.tensor([1.0, 0.0, 1.0, 0.0]) loss = loss_fn(predictions, targets) print(loss)

Результат выполнения кода:

tensor(0.4016)

Пример с разными режимами агрегации

Рассмотрим использование различных режимов агрегации потерь:

import torch import torch.nn as nn predictions = torch.tensor([0.9, 0.2, 0.7, 0.4]) targets = torch.tensor([1.0, 0.0, 1.0, 0.0]) loss_none = nn.BCELoss(reduction='none') loss_mean = nn.BCELoss(reduction='mean') loss_sum = nn.BCELoss(reduction='sum') print(loss_none(predictions, targets)) print(loss_mean(predictions, targets)) print(loss_sum(predictions, targets))

Результат выполнения кода:

tensor([0.1054, 0.2231, 0.3567, 0.5108]) tensor(0.2990) tensor(1.1960)

Пример с учетом дисбаланса классов

Используем параметр pos_weight для борьбы с дисбалансом классов:

import torch import torch.nn as nn pos_weight = torch.tensor([3.0]) loss_fn = nn.BCELoss(pos_weight=pos_weight) predictions = torch.tensor([0.3, 0.8, 0.6]) targets = torch.tensor([1.0, 1.0, 0.0]) loss = loss_fn(predictions, targets) print(loss)

Результат выполнения кода:

tensor(0.7641)

Пример использования в обучении

Покажем, как использовать BCELoss в цикле обучения модели:

import torch import torch.nn as nn torch.manual_seed(0) model = nn.Linear(5, 1) loss_fn = nn.BCELoss() optimizer = torch.optim.SGD(model.parameters(), lr=0.01) inputs = torch.randn(4, 5) targets = torch.tensor([1.0, 0.0, 1.0, 0.0]) outputs = torch.sigmoid(model(inputs)) loss = loss_fn(outputs.squeeze(), targets) optimizer.zero_grad() loss.backward() optimizer.step() print(f"Loss: {loss.item():.4f}")

Результат выполнения кода:

"Loss: 0.7747"

Пример с батчами

Рассмотрим вычисление потерь для пакета данных большего размера:

import torch import torch.nn as nn torch.manual_seed(0) loss_fn = nn.BCELoss() batch_size = 8 num_features = 3 predictions = torch.rand(batch_size, num_features) targets = torch.randint(0, 2, (batch_size, num_features)).float() loss = loss_fn(predictions, targets) print(f"Batch loss: {loss.item():.4f}") print(f"Predictions shape: {predictions.shape}") print(f"Targets shape: {targets.shape}")

Результат выполнения кода:

"Batch loss: 0.8897" "Predictions shape: torch.Size([8, 3])" "Targets shape: torch.Size([8, 3])"

Смотрите также

  • класс BCEWithLogitsLoss,
    который объединяет сигмоиду и бинарную кросс-энтропию для численной стабильности
  • класс CrossEntropyLoss,
    который используется для задач многоклассовой классификации
  • класс MSELoss,
    который вычисляет среднеквадратичную ошибку для задач регрессии
  • класс NLLLoss,
    который реализует отрицательную логарифмическую вероятность
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить