PyTorch torch.nn.BCEWithLogitsLoss Function
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 weightspos_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())
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())
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())
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.
Other Extensions