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
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 FunctionPrediction (input) ShapeLabel (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
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

# Label smoothing, alleviates overfitting, often used in image classification competitions
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

criterion = nn.BCELoss()

# 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

criterion = nn.BCEWithLogitsLoss()

# 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

# pos_weight: weight for positive samples; the larger the value, the more attention paid to positive samples
# 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

criterion = nn.NLLLoss()

# 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 usingCrossEntropyLoss。


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

criterion = nn.MSELoss()

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

criterion = nn.L1Loss()

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

criterion = nn.SmoothL1Loss()

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
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

# delta controls the switching point between MSE and L1
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

criterion = nn.KLDivLoss(reduction='batchmean')

# 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

temperature = 4.0  # Temperature coefficient, the larger the softer

# 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

criterion = nn.MarginRankingLoss(margin=0.5)

# 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

criterion = nn.TripletMarginLoss(margin=1.0)

# 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

criterion = nn.CTCLoss(blank=0)  # blank label index

# 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

predictions = torch.tensor([1.0, 2.0, 3.0, 4.0])
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

# Give higher weight to samples with larger errors (focal loss idea)
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

# Assume 3 classes: class 0 has 1000, class 1 has 100, class 2 has 50
# 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

# ignore_index: ignore this label when computing the loss
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
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
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 typeRecommended loss functionNotes
Multi-class classificationCrossEntropyLossMost 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 classificationBCEWithLogitsLossMore stable than BCELoss
Multi-label classificationBCEWithLogitsLossEach label judged independently
Object detection (classification head)CrossEntropyLoss / Focal LossUse Focal when positive-negative samples are imbalanced
Object detection (regression head)SmoothL1Loss / GIoULossStandard practice
Ordinary regressionMSELossFirst choice when there are no outliers
Regression with outliersHuberLoss / SmoothL1LossRobust regression
Image segmentationCrossEntropyLoss + DiceLossCombined use works better
Speech recognitionCTCLossSequence alignment
Metric learning / face recognitionTripletMarginLossLearn feature space distance
Knowledge distillationKLDivLossLearn soft label distribution

Common Misuses and Notes

Example

# Wrong: manually applied Softmax before CrossEntropyLoss
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
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}")
Other Extensions