PyTorch Loss Functions
The loss function measures the gap between the model's predictions and the true values, and serves as the core guide for neural network training—the optimizer updates model parameters by minimizing the loss function.
PyTorch, in itstorch.nnmodule, has built-in more than ten common loss functions, covering major task types such as classification, regression, and ranking.
1. Loss Function Basics
Basic Usage
All PyTorch loss functions arenn.Modulesubclasses of , with unified usage:
Example
import torch.nn as nn
# 1. Instantiate the loss function
criterion = nn.CrossEntropyLoss()
# 2. Compute the loss (predictions first, ground truth second)
loss = criterion(predictions, targets)
# 3. Backpropagation
loss.backward()
Shape Conventions for Prediction Values
Different loss functions have different shape requirements for inputs; this is the most common source of errors for beginners:
| Loss Function | Prediction (input) Shape | Label (target) Shape |
|---|---|---|
CrossEntropyLoss | (N, C)Raw logits | (N,)Integer class indices |
BCELoss | (N,)Probabilities after Sigmoid | (N,)0/1 floating-point numbers |
BCEWithLogitsLoss | (N,)Raw logits | (N,)0/1 floating-point numbers |
MSELoss | (N,)Any real number | (N,)Any real number |
NLLLoss | (N, C)Probabilities after log_softmax | (N,)Integer class indices |
N = batch size,C= number of classes
2. Classification Task Loss Functions
2.1 CrossEntropyLoss (Cross-Entropy Loss)
The most commonly used multi-class loss function,internally performs Softmax + log + negative sign automatically; no need to manually apply Softmax to the model output.
Mathematical formula:
Loss = -sum(y_c * log(p_c))
where p_c = exp(x_c) / sum_j exp(x_j) is the Softmax output.
Example
import torch.nn as nn
criterion = nn.CrossEntropyLoss()
# Model output: raw logits, shape (batch_size, num_classes)
# No need to apply Softmax in advance!
predictions = torch.tensor([
[2.0, 0.5, 0.3], # Sample 1, most likely class 0
[0.1, 3.0, 0.2], # Sample 2, most likely class 1
[0.2, 0.1, 4.0], # Sample 3, most likely class 2
])
# Labels: integer class indices, shape (batch_size,)
targets = torch.tensor([0, 1, 2])
loss = criterion(predictions, targets)
print(f"Loss: {loss.item():.4f}") # Loss: 0.1763
Supports soft labels (Label Smoothing):
Example
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
# Also supports directly passing soft labels (probability distributions)
soft_targets = torch.tensor([
[0.9, 0.05, 0.05],
[0.05, 0.9, 0.05],
])
predictions = torch.randn(2, 3)
loss = criterion(predictions, soft_targets)
Applicable scenarios:All multi-class tasks such as multi-class classification (cat/dog/bird), image classification, text classification, etc.
2.2 BCELoss (Binary Cross-Entropy Loss)
Specifically used forbinary classificationorand multi-label classificationtasks. The input must beSigmoidprobability values after processing (0-1).
Mathematical formula:
Loss = -[y * log(p) + (1-y) * log(1-p)]
Example
# The model output must first pass through Sigmoid, range (0, 1)
raw_output = torch.tensor([2.0, -1.0, 0.5, -3.0])
predictions = torch.sigmoid(raw_output) # [0.88, 0.27, 0.62, 0.05]
# Labels: float type 0.0 or 1.0
targets = torch.tensor([1.0, 0.0, 1.0, 0.0])
loss = criterion(predictions, targets)
print(f"Loss: {loss.item():.4f}") # Loss: 0.2824
# Multi-label classification (each sample can belong to multiple classes)
# predictions shape: (batch_size, num_labels)
predictions_ml = torch.sigmoid(torch.randn(4, 5))
targets_ml = torch.randint(0, 2, (4, 5)).float()
loss_ml = criterion(predictions_ml, targets_ml)
BCELossThe input is required to be in the range (0, 1). Passing raw logits will cause numerical instability or even NaN. It is recommended to use the improved version below,BCEWithLogitsLoss。
2.3 BCEWithLogitsLoss
BCELosswhich is an improved version of ,and automatically performs Sigmoid internally, providing better numerical stability and recommended as the first choice.
Example
# Pass raw logits directly, no need for manual Sigmoid
predictions = torch.tensor([2.0, -1.0, 0.5, -3.0])
targets = torch.tensor([1.0, 0.0, 1.0, 0.0])
loss = criterion(predictions, targets)
print(f"Loss: {loss.item():.4f}")
# Equivalent to (but with better numerical stability):
# loss = BCELoss(Sigmoid(predictions), targets)
With positive sample weights (handling class imbalance):
Example
# For example, if negative samples are 10 times positive samples, set pos_weight=10
pos_weight = torch.tensor([10.0])
criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)
Applicable scenarios:Binary classification (spam detection), multi-label classification (multi-label tagging of articles), object detection (foreground/background judgment).
2.4 NLLLoss (Negative Log-Likelihood Loss)
Requires manually applying to the model outputlog_softmax, which provides higher flexibility.CrossEntropyLoss = LogSoftmax + NLLLoss。
Example
# Must manually apply log_softmax first
raw_output = torch.randn(4, 3) # (batch, num_classes)
log_probs = torch.log_softmax(raw_output, dim=1)
targets = torch.tensor([0, 2, 1, 0])
loss = criterion(log_probs, targets)
Use cases:When log probabilities are needed in intermediate steps (e.g., CTC, Beam Search); in other cases, prefer using
CrossEntropyLoss。
3. Regression Task Loss Functions
3.1 MSELoss (Mean Squared Error)
The most classic regression loss, which isvery sensitive to large errors(because squaring amplifies the impact of large errors).
Mathematical formula:
MSELoss = (1/N) * sum((y_i - y_hat_i)^2)
Example
predictions = torch.tensor([2.5, 0.5, 2.0, 8.0])
targets = torch.tensor([3.0, -0.5, 2.0, 7.0])
loss = criterion(predictions, targets)
print(f"MSE Loss: {loss.item():.4f}") # MSE Loss: 0.3750
# Manual verification
manual = ((predictions - targets) ** 2).mean()
print(f"Manual calculation: {manual.item():.4f}") # 0.3750
Applicable scenarios:Continuous value regression such as house price prediction and temperature prediction; works well when there are no obvious outliers in the data.
3.2 L1Loss (Mean Absolute Error)
PairMore robust to outliersbecause it uses absolute value instead of squaring, so large errors are not over-amplified.
Mathematical formula:
L1Loss = (1/N) * sum(|y_i - y_hat_i|)
Example
predictions = torch.tensor([2.5, 0.5, 2.0, 8.0])
targets = torch.tensor([3.0, -0.5, 2.0, 7.0])
loss = criterion(predictions, targets)
print(f"L1 Loss: {loss.item():.4f}") # L1 Loss: 0.5000
3.3 SmoothL1Loss (Huber Loss)
Combines the advantages of MSE and L1: uses MSE for small errors (smooth, stable gradients), and L1 for large errors (outlier resistant). It is the standard loss in object detection (Faster R-CNN).
Mathematical formula:
SmoothL1(x) = 0.5*x^2 if |x| < 1, else |x| - 0.5
Example
predictions = torch.tensor([2.5, 0.5, 2.0, 8.0])
targets = torch.tensor([3.0, -0.5, 2.0, 7.0])
loss = criterion(predictions, targets)
print(f"SmoothL1 Loss: {loss.item():.4f}")
Comparison of the three regression losses:
Example
import torch.nn as nn
predictions = torch.tensor([0.0, 1.0, 5.0, 10.0]) # Simulate errors of different magnitudes
targets = torch.zeros(4)
for name, fn in [("MSELoss", nn.MSELoss(reduction='none')),
("L1Loss", nn.L1Loss(reduction='none')),
("SmoothL1",nn.SmoothL1Loss(reduction='none'))]:
losses = fn(predictions, targets)
print(f"{name:12s}: {[f'{l:.2f}' for l in losses.tolist()]}")
输出示例: MSELoss : ['0.00', '1.00', '25.00', '100.00'] # 大误差被平方放大 L1Loss : ['0.00', '1.00', '5.00', '10.00'] # 线性增长 SmoothL1 : ['0.00', '0.50', '4.50', '9.50'] # 中间值
Applicable scenarios:Regression tasks with both small errors and outliers, such as bounding box regression in object detection and depth estimation.
4. Advanced Loss Functions
4.1 HuberLoss
SmoothL1LossA generalized version of , allowing custom threshold switchingdelta(default 1.0).
Example
criterion = nn.HuberLoss(delta=1.5)
predictions = torch.randn(10)
targets = torch.randn(10)
loss = criterion(predictions, targets)
4.2 KLDivLoss (KL Divergence)
Measures the difference between two probability distributions, commonly used inknowledge distillationandvariational autoencoder (VAE)。
Mathematical formula:
KL(P || Q) = sum(P(i) * log(P(i) / Q(i)))
Example
# input must be log probabilities, target is ordinary probabilities
log_predictions = torch.log_softmax(torch.randn(4, 5), dim=1)
targets = torch.softmax(torch.randn(4, 5), dim=1)
loss = criterion(log_predictions, targets)
print(f"KL Div Loss: {loss.item():.4f}")
Typical usage in knowledge distillation:
Example
# Teacher model output
teacher_logits = torch.randn(32, 10)
# Student model output
student_logits = torch.randn(32, 10)
soft_labels = torch.softmax(teacher_logits / temperature, dim=1)
soft_preds = torch.log_softmax(student_logits / temperature, dim=1)
distill_loss = nn.KLDivLoss(reduction='batchmean')(soft_preds, soft_labels)
distill_loss *= temperature ** 2 # Restore gradient magnitude
4.3 MarginRankingLoss (Ranking Loss)
Determines the relative order of two inputs, commonly used inlearning to rankandsimilarity learning。
Example
# x1 should be closer to target than x2 (y=1 means x1 > x2)
x1 = torch.tensor([0.8, 0.3, 0.6])
x2 = torch.tensor([0.2, 0.7, 0.5])
y = torch.tensor([1.0, -1.0, 1.0]) # 1: x1>x2, -1: x1<x2
loss = criterion(x1, x2, y)
4.4 TripletMarginLoss (Triplet Loss)
Used for metric learning, requiringanchorandpositive (same class)the distance to be less than that withnegative (different class)distance.
Example
# Each vector has dimension embedding_dim
anchor = torch.randn(32, 128) # Anchor sample
positive = torch.randn(32, 128) # Positive sample (same class)
negative = torch.randn(32, 128) # Negative sample (different class)
loss = criterion(anchor, positive, negative)
# Objective: dist(anchor, positive) + margin < dist(anchor, negative)
Applicable scenarios:Face recognition, image retrieval, few-shot learning.
4.5 CTCLoss (Sequence Labeling Loss)
Used forsequence tasks where input and output lengths are misaligned,such as speech recognition (acoustic sequence -> text sequence), handwriting recognition.
Example
# log_probs: (T, N, C) T=time steps, N=batch, C=number of classes
T, N, C = 50, 4, 20
log_probs = torch.log_softmax(torch.randn(T, N, C), dim=2)
targets = torch.randint(1, C, (N * 10,)) # Concatenated target sequence
input_lengths = torch.full((N,), T, dtype=torch.long)
target_lengths = torch.full((N,), 10, dtype=torch.long)
loss = criterion(log_probs, targets, input_lengths, target_lengths)
5. Detailed Explanation of the reduction Parameter
All loss functions support thereductionparameter, which controls how sample losses are aggregated:
Example
targets = torch.tensor([1.5, 2.5, 2.0, 5.0])
# Per-sample errors: [0.25, 0.25, 1.00, 1.00]
# mean (default): average over all samples
loss_mean = nn.MSELoss(reduction='mean')(predictions, targets)
print(f"mean: {loss_mean.item():.4f}") # 0.6250
# sum: sum over all samples
loss_sum = nn.MSELoss(reduction='sum')(predictions, targets)
print(f"sum: {loss_sum.item():.4f}") # 2.5000
# none: returns each sample's individual loss (often used for weighting)
loss_none = nn.MSELoss(reduction='none')(predictions, targets)
print(f"none: {loss_none.tolist()}") # [0.25, 0.25, 1.0, 1.0]
reduction='none'Practical application — weighting different samples:
Example
per_sample_loss = nn.MSELoss(reduction='none')(predictions, targets)
weights = torch.tensor([1.0, 1.0, 2.0, 2.0]) # Manually set weights
weighted_loss = (per_sample_loss * weights).mean()
6. Class Weights and Sample Weights
Class Weights (Handling Class Imbalance)
When certain classes have very few samples in the dataset, give minority classes higher weights:
Example
# Weight inversely proportional to frequency
class_counts = torch.tensor([1000.0, 100.0, 50.0])
weights = 1.0 / class_counts
weights = weights / weights.sum() * len(weights) # Normalize
criterion = nn.CrossEntropyLoss(weight=weights)
Ignoring Specific Labels
In tasks such as semantic segmentation, it is often necessary to ignore boundary pixels (labeled 255):
Example
criterion = nn.CrossEntropyLoss(ignore_index=255)
# Semantic segmentation scenario
predictions = torch.randn(2, 21, 256, 256) # (N, C, H, W)
targets = torch.randint(0, 22, (2, 256, 256))
targets[targets == 21] = 255 # Boundaries labeled as 255
loss = criterion(predictions, targets)
7. Custom Loss Functions
When built-in loss functions cannot meet the requirements, you can customize them in two ways:
Method 1: Functional (Simple)
Example
import torch.nn.functional as F
def focal_loss(predictions, targets, gamma=2.0, alpha=0.25):
"""
Focal Loss: solves the severe positive-negative sample imbalance problem in object detection
Reduces the weight of easily classified samples, allowing the model to focus on hard samples
"""
ce_loss = F.cross_entropy(predictions, targets, reduction='none')
pt = torch.exp(-ce_loss) # Probability of correct prediction
focal_weight = alpha * (1 - pt) ** gamma # Hard samples get higher weight
return (focal_weight * ce_loss).mean()
# Use
predictions = torch.randn(8, 10)
targets = torch.randint(0, 10, (8,))
loss = focal_loss(predictions, targets)
Method 2: Subclassing nn.Module (Recommended)
Example
import torch.nn as nn
class DiceLoss(nn.Module):
"""
Dice Loss: commonly used in image segmentation, directly optimizes the Dice coefficient
More robust than CrossEntropy to class imbalance (e.g., small object segmentation)
"""
def __init__(self, smooth=1.0):
super().__init__()
self.smooth = smooth
def forward(self, predictions, targets):
# predictions: (N, C, H, W) -> probabilities after sigmoid
# targets: (N, C, H, W) -> one-hot encoded labels
predictions = torch.sigmoid(predictions)
# Flatten to (N, -1)
pred_flat = predictions.view(predictions.size(0), -1)
target_flat = targets.view(targets.size(0), -1).float()
intersection = (pred_flat * target_flat).sum(dim=1)
dice = (2.0 * intersection + self.smooth) / (
pred_flat.sum(dim=1) + target_flat.sum(dim=1) + self.smooth
)
return 1 - dice.mean()
class CombinedLoss(nn.Module):
"""
Combined loss: CrossEntropy + Dice, balancing pixel-level classification and region overlap
Common combination for image segmentation
"""
def __init__(self, ce_weight=0.5, dice_weight=0.5):
super().__init__()
self.ce_weight = ce_weight
self.dice_weight = dice_weight
self.ce = nn.CrossEntropyLoss()
self.dice = DiceLoss()
def forward(self, predictions, targets):
return (self.ce_weight * self.ce(predictions, targets) +
self.dice_weight * self.dice(predictions, targets))
# Use
criterion = CombinedLoss(ce_weight=0.4, dice_weight=0.6)
8. Loss Function Selection Guide
Select by Task Type
| Task type | Recommended loss function | Notes |
|---|---|---|
| Multi-class classification | CrossEntropyLoss | Most general, preferred choice |
| Multi-class classification (class imbalance) | CrossEntropyLoss(weight=...) | Weight minority classes |
| Multi-class classification (noisy labels) | CrossEntropyLoss(label_smoothing=0.1) | Prevents overfitting |
| Binary classification | BCEWithLogitsLoss | More stable than BCELoss |
| Multi-label classification | BCEWithLogitsLoss | Each label judged independently |
| Object detection (classification head) | CrossEntropyLoss / Focal Loss | Use Focal when positive-negative samples are imbalanced |
| Object detection (regression head) | SmoothL1Loss / GIoULoss | Standard practice |
| Ordinary regression | MSELoss | First choice when there are no outliers |
| Regression with outliers | HuberLoss / SmoothL1Loss | Robust regression |
| Image segmentation | CrossEntropyLoss + DiceLoss | Combined use works better |
| Speech recognition | CTCLoss | Sequence alignment |
| Metric learning / face recognition | TripletMarginLoss | Learn feature space distance |
| Knowledge distillation | KLDivLoss | Learn soft label distribution |
Common Misuses and Notes
Example
output = torch.softmax(model(x), dim=1) # Redundant softmax
loss = nn.CrossEntropyLoss()(output, targets) # It will be applied again internally
# Correct: pass logits directly
output = model(x) # Raw logits
loss = nn.CrossEntropyLoss()(output, targets)
# Wrong: passing raw values without sigmoid to BCELoss
loss = nn.BCELoss()(model(x), targets) # May exceed [0,1], numerically unstable
# Correct: use BCEWithLogitsLoss
loss = nn.BCEWithLogitsLoss()(model(x), targets)
# Wrong: label type mismatch (integer vs float)
targets = torch.tensor([1, 0, 1]) # Integer type
loss = nn.BCEWithLogitsLoss()(preds, targets) # Error!
# Correct: BCEWithLogitsLoss requires float labels
targets = torch.tensor([1.0, 0.0, 1.0]) # Float type
loss = nn.BCEWithLogitsLoss()(preds, targets) # Correct
# Wrong: loss not extracted with .item(), causing the computation graph to keep accumulating and GPU memory to overflow
total_loss += loss # loss is a tensor and holds the computation graph
# Correct: use .item() to extract a scalar
total_loss += loss.item()
Full Training Example
Example
import torch.nn as nn
import torch.optim as optim
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Model & loss function & optimizer
model = MyModel().to(device)
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
for epoch in range(num_epochs):
model.train()
total_loss, correct = 0.0, 0
for inputs, labels in train_loader:
inputs = inputs.to(device)
labels = labels.to(device)
optimizer.zero_grad()
outputs = model(inputs) # Raw logits
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * inputs.size(0) # .item() extracts a scalar
correct += (outputs.argmax(1) == labels).sum().item()
avg_loss = total_loss / len(train_loader.dataset)
accuracy = correct / len(train_loader.dataset)
print(f"Epoch {epoch+1} | Loss: {avg_loss:.4f} | Acc: {accuracy:.4f}")