PyTorch torch.nn.CrossEntropyLoss Function

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


torch.nn.CrossEntropyLossis a loss function used for multi-class classification in PyTorch.

It combines nn.LogSoftmax and nn.NLLLoss, commonly used in tasks such as image classification and text classification.

Function Definition

torch.nn.CrossEntropyLoss(weight=None, ignore_index=-100, reduction='mean', label_smoothing=0.0)

Parameter Description:

  • weight(Tensor): Assigns different weights to each class, used for class imbalance situations.
  • ignore_index(int): Ignores the loss calculation for the specified index. Defaults to -100.
  • reduction(str): Loss reduction method. Options are'mean'、'sum'、'none'. Defaults to'mean'。
  • label_smoothing(float): Label smoothing parameter, ranging from 0 to 1. Defaults to 0.

Mathematical Principle

The formula for cross-entropy loss:

Loss = -log(exp(y_true) / sum(exp(y_i)))

That is, the larger the predicted probability of the correct class, the smaller the loss.


Usage Examples

Example 1: Basic Usage

Create and use cross-entropy loss:

Example

import torch
import torch.nn as nn

# Create loss function
criterion = nn.CrossEntropyLoss()

# Logits output by the model (unnormalized)
# Shape: (batch_size, num_classes)
outputs = torch.randn(4, 10)

# True labels
labels = torch.tensor([2, 5, 1, 7])

# Compute loss
loss = criterion(outputs, labels)

print("Model output (logits):", outputs[0].tolist())
print("True labels:", labels[0].item())
print("Cross-entropy loss:", loss.item())

Example 2: Class Weights

Handling class imbalance:

Example

import torch
import torch.nn as nn

# Class weights: give higher weight to minority classes
# Assume 10 classes, class 3 and class 7 are more important
weight = torch.ones(10)
weight[3] = 2.0
weight[7] = 2.0

criterion_weighted = nn.CrossEntropyLoss(weight=weight)

outputs = torch.randn(4, 10)
labels = torch.tensor([2, 3, 7, 5])

loss = criterion_weighted(outputs, labels)
print("Weighted cross-entropy loss:", loss.item())

Example 3: Label Smoothing

Use label smoothing to prevent overfitting:

Example

import torch
import torch.nn as nn

# Label smoothing: 0.1 means distributing 10% of the probability uniformly to other classes
criterion_smooth = nn.CrossEntropyLoss(label_smoothing=0.1)

outputs = torch.randn(4, 10)
labels = torch.tensor([2, 5, 1, 7])

loss = criterion_smooth(outputs, labels)
print("Loss with label smoothing:", loss.item())

# Comparison: without label smoothing
criterion = nn.CrossEntropyLoss()
loss_no_smooth = criterion(outputs, labels)
print("Loss without label smoothing:", loss_no_smooth.item())

Example 4: Complete Classification Training Pipeline

A complete model training example:

Example

import torch
import torch.nn as nn
import torch.optim as optim

# Simple classification model
class Classifier(nn.Module):
    def __init__(self, input_dim=784, num_classes=10):
        super(Classifier, self).__init__()
        self.fc = nn.Sequential(
            nn.Linear(input_dim, 256),
            nn.ReLU(),
            nn.Dropout(0.5),
            nn.Linear(256, num_classes)
        )

    def forward(self, x):
        return self.fc(x)

# Initialize model and loss
model = Classifier()
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

# Simulate training data
batch_size = 32
x = torch.randn(batch_size, 784)  # Input
y = torch.randint(0, 10, (batch_size,))  # Labels

# Forward pass
model.train()
outputs = model(x)
loss = criterion(outputs, y)

# Backward pass
optimizer.zero_grad()
loss.backward()
optimizer.step()

print("Batch loss:", loss.item())

# Prediction
model.eval()
with torch.no_grad():
    outputs = model(x)
    predictions = outputs.argmax(dim=1)
    accuracy = (predictions == y).float().mean()

print("Prediction accuracy:", accuracy.item())

Example 5: Using ignore_index

Ignore specific labels:

Example

import torch
import torch.nn as nn

# Ignore samples with label=-100
criterion = nn.CrossEntropyLoss(ignore_index=-100)

outputs = torch.randn(5, 10)
# Some samples have label -100, meaning they are ignored
labels = torch.tensor([2, -100, 5, -100, 7])

loss = criterion(outputs, labels)
print("Loss after ignoring special labels:", loss.item())

Example 6: Different Reduction Modes

Control the loss reduction method:

Example

import torch
import torch.nn as nn

outputs = torch.randn(4, 10)
labels = torch.tensor([2, 5, 1, 7])

# mean: returns the average loss
loss_mean = nn.CrossEntropyLoss(reduction='mean')(outputs, labels)
print("mean:", loss_mean.item())

# sum: returns the total sum
loss_sum = nn.CrossEntropyLoss(reduction='sum')(outputs, labels)
print("sum:", loss_sum.item())

# none: returns the loss for each sample
loss_none = nn.CrossEntropyLoss(reduction='none')(outputs, labels)
print("none:", loss_none.tolist())

Frequently Asked Questions

Q1: What is the difference between CrossEntropyLoss and NLLLoss?

CrossEntropyLoss = LogSoftmax + NLLLoss. It already has softmax built in, so there is no need to add it manually.

Q2: Why doesn't the model output use softmax?

CrossEntropyLoss automatically computes softmax internally, and using logits directly improves numerical stability.

Q3: What scenarios is label smoothing suitable for?

Label smoothing is suitable for situations with a large number of classes, and it can improve the model's generalization ability.


Use Cases

nn.CrossEntropyLossThe main application scenarios include:

  • Image classification: such as CIFAR-10, ImageNet
  • Text classification: sentiment analysis, topic classification
  • Multi-class classification tasks: any classification task with more than 2 classes

Note: Labels should be class indices (0 to num_classes-1), not one-hot encoded.


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

Other Extensions