PyTorch torch.nn.CrossEntropyLoss Function
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.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.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.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.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.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.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.
Other Extensions