PyTorch torch.nn.BCEWithLogitsLoss Function

PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual


torch.nn.BCEWithLogitsLossIt is the binary cross-entropy loss function (with Sigmoid) in PyTorch.

It combines Sigmoid and BCE, making it numerically more stable, used for binary classification tasks.

Function Definition

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

Parameters:

  • weight: Manual weights
  • pos_weight: Positive class weight, used for class imbalance

Usage Examples

Example 1: Basic Usage

Example

import torch
import torch.nn as nn

criterion = nn.BCEWithLogitsLoss()

# Unnormalized logits
logits = torch.tensor([2.0, -1.0, 0.5, -3.0])
# Binary labels
targets = torch.tensor([1.0, 0.0, 1.0, 0.0])

loss = criterion(logits, targets)
print("BCE Loss:", loss.item())

# Manual verification: Sigmoid + BCE
sigmoid = torch.sigmoid(logits)
bce = nn.BCELoss()(sigmoid, targets)
print("Manual BCE:", bce.item())

Example 2: Class Imbalance

Example

import torch
import torch.nn as nn

# Positive class weight: increase the importance of positive samples
pos_weight = torch.tensor([5.0])
criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)

logits = torch.randn(10, 1)
targets = torch.zeros(10, 1)
targets[:2] = 1.0  # Positive class is rare

loss = criterion(logits, targets)
print("Weighted BCE Loss:", loss.item())

Example 3: Multi-label Classification

Example

import torch
import torch.nn as nn

# Multi-label binary classification
criterion = nn.BCEWithLogitsLoss()

# batch=4, 5 classes, each can be 0 or 1
logits = torch.randn(4, 5)
labels = torch.randint(0, 2, (4, 5)).float()

loss = criterion(logits, labels)
print("Multi-label Loss:", loss.item())

Use Cases

  • Binary classification: Single label
  • Multi-label classification: Each label is independent
  • Class imbalance: Use pos_weight

Note: Input is logits, no need to apply Sigmoid beforehand.


PyTorch torch.nn 参考手册PyTorch torch.nn Reference Manual

Other Extensions