别再死记公式了!用PyTorch和TensorFlow实战理解交叉熵损失函数
从代码实践反推交叉熵:PyTorch/TensorFlow中的损失函数可视化教学
在咖啡厅里,我遇到过一位盯着屏幕发呆的学员,他面前摊开的笔记本上写满了交叉熵的数学公式。"我背了三天公式,可还是不知道这玩意儿在模型里到底怎么工作的。"他苦恼地说。这场景让我意识到,太多机器学习教程把交叉熵讲成了抽象数学概念,而忽略了它作为损失函数最本质的实践价值。本文将带你用PyTorch和TensorFlow的代码实操,从反向理解这个深度学习中最重要的损失函数之一。
1. 为什么我们需要交叉熵损失?
想象你正在训练一个猫狗分类器。当模型预测一张明显是狗的照片有80%概率是猫时,我们该如何用数字量化这个错误的"离谱程度"?这就是交叉熵的用武之地——它测量的是预测概率分布与真实分布的"距离"。
与MSE(均方误差)不同,交叉熵特别适合分类问题,因为它:
- 对错误预测惩罚更严厉:预测狗为猫的概率从50%提升到70%,交叉熵的惩罚增长比MSE更显著
- 与softmax天然契合:配合softmax输出层时,梯度计算更加高效
- 信息论基础:本质上衡量的是用预测分布编码真实分布所需的额外比特数
# 直观对比MSE和交叉熵对错误预测的敏感度
import numpy as np
def mse_loss(y_true, y_pred):
return np.mean((y_true - y_pred)**2)
def cross_entropy_loss(y_true, y_pred):
return -np.sum(y_true * np.log(y_pred))
# 真实标签:类别2(one-hot编码)
y_true = np.array([0, 0, 1])
# 两个预测案例
y_pred_1 = np.array([0.2, 0.3, 0.5]) # 勉强正确
y_pred_2 = np.array([0.8, 0.1, 0.1]) # 严重错误
print(f"MSE差异:{mse_loss(y_true, y_pred_2)/mse_loss(y_true, y_pred_1):.1f}x")
print(f"交叉熵差异:{cross_entropy_loss(y_true, y_pred_2)/cross_entropy_loss(y_true, y_pred_1):.1f}x")
执行这段代码你会看到,当预测从"勉强正确"变为"严重错误"时,MSE只增长了约2倍,而交叉熵增长了超过8倍——这正是分类任务需要的特性。
2. 框架中的交叉熵实现细节
2.1 PyTorch的nn.CrossEntropyLoss
PyTorch的实现有几个关键特性常被忽略:
- 内置softmax:实际上不需要在模型最后一层添加softmax
- 输入格式:接受原始logits(未归一化的分数)而非概率
- 标签处理:可以使用类索引(如2)而非one-hot编码
import torch
import torch.nn as nn
# 模拟一个batch_size=3,类别数=5的分类任务
logits = torch.randn(3, 5) # 原始网络输出
labels = torch.tensor([1, 0, 4]) # 每个样本的真实类别索引
loss_fn = nn.CrossEntropyLoss()
loss = loss_fn(logits, labels)
print(f"计算得到的损失值:{loss.item():.4f}")
# 手动验证计算过程
def manual_ce(logits, labels):
# 第一步:对logits应用softmax
probabilities = torch.softmax(logits, dim=1)
# 第二步:取对应真实类别的概率
true_class_probs = probabilities[torch.arange(len(labels)), labels]
# 第三步:计算负对数
return -torch.mean(torch.log(true_class_probs))
manual_loss = manual_ce(logits, labels)
print(f"手动计算结果:{manual_loss.item():.4f}")
注意:PyTorch的实现做了数值稳定优化,实际代码中会使用log_softmax和nll_loss的组合,避免直接计算可能出现的数值问题。
2.2 TensorFlow的CategoricalCrossentropy
TensorFlow提供了更灵活的选择:
import tensorflow as tf
# 情况1:输入已经是概率(配合softmax输出层)
loss_fn_prob = tf.keras.losses.CategoricalCrossentropy(from_logits=False)
# 情况2:输入是logits(配合线性输出层)
loss_fn_logits = tf.keras.losses.CategoricalCrossentropy(from_logits=True)
# 模拟数据
y_true = tf.constant([[0, 0, 1], [1, 0, 0], [0, 1, 0]]) # one-hot
y_pred_logits = tf.random.normal((3, 3))
y_pred_probs = tf.nn.softmax(y_pred_logits)
# 两种计算方式结果应当接近
loss1 = loss_fn_prob(y_true, y_pred_probs)
loss2 = loss_fn_logits(y_true, y_pred_logits)
print(f"概率输入损失:{loss1.numpy():.4f}")
print(f"logits输入损失:{loss2.numpy():.4f}")
关键区别:
from_logits=True时,内部会自动应用softmax,通常更数值稳定from_logits=False时,需确保输入已经是合法的概率分布(每行和为1)
3. 可视化交叉熵的训练动态
理解损失函数最好的方式就是观察它在真实训练中的表现。我们用CIFAR-10数据集做个实验:
import matplotlib.pyplot as plt
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
# 准备CIFAR-10数据
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
train_set = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
# 简单模型
model = torch.nn.Sequential(
torch.nn.Flatten(),
torch.nn.Linear(32*32*3, 512),
torch.nn.ReLU(),
torch.nn.Linear(512, 10) # 输出logits
)
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
loss_fn = torch.nn.CrossEntropyLoss()
# 训练循环
loss_history = []
for epoch in range(5):
epoch_loss = 0
for images, labels in train_loader:
optimizer.zero_grad()
outputs = model(images)
loss = loss_fn(outputs, labels)
loss.backward()
optimizer.step()
epoch_loss += loss.item()
avg_loss = epoch_loss / len(train_loader)
loss_history.append(avg_loss)
print(f"Epoch {epoch+1}, Loss: {avg_loss:.4f}")
# 绘制损失曲线
plt.plot(loss_history)
plt.xlabel('Epoch')
plt.ylabel('Cross Entropy Loss')
plt.title('Training Loss on CIFAR-10')
plt.show()
运行这段代码,你会看到典型的损失下降曲线。但更有趣的是观察单个batch的损失变化:
# 获取第一个batch的数据
first_batch = next(iter(train_loader))
images, labels = first_batch
# 在训练前后分别计算这个batch的损失
with torch.no_grad():
initial_outputs = model(images)
initial_loss = loss_fn(initial_outputs, labels)
# 训练几个step后...
for _ in range(100):
optimizer.zero_grad()
outputs = model(images)
loss = loss_fn(outputs, labels)
loss.backward()
optimizer.step()
with torch.no_grad():
trained_outputs = model(images)
trained_loss = loss_fn(trained_outputs, labels)
print(f"初始损失:{initial_loss.item():.4f}")
print(f"训练后损失:{trained_loss.item():.4f}")
# 可视化预测变化
def plot_probs(output, title):
probs = torch.softmax(output, dim=1)[0]
plt.bar(range(10), probs.detach().numpy())
plt.xlabel('Class')
plt.ylabel('Probability')
plt.title(title)
plt.xticks(range(10), ['plane', 'car', 'bird', 'cat', 'deer',
'dog', 'frog', 'horse', 'ship', 'truck'])
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plot_probs(initial_outputs, "Initial Predictions")
plt.subplot(1, 2, 2)
plot_probs(trained_outputs, "After 100 Steps")
plt.show()
这个可视化展示了交叉熵如何驱动模型调整预测——从初始的均匀分布逐渐向真实标签集中。
4. 常见陷阱与调试技巧
4.1 输入尺度问题
交叉熵对输入logits的尺度非常敏感:
# 模拟不同尺度的logits
small_logits = torch.tensor([[1.0, 2.0, 3.0]])
large_logits = torch.tensor([[100.0, 200.0, 300.0]])
labels = torch.tensor([2])
small_loss = loss_fn(small_logits, labels)
large_loss = loss_fn(large_logits, labels)
print(f"小尺度logits损失:{small_loss.item():.4f}")
print(f"大尺度logits损失:{large_loss.item():.4f}")
# 查看对应的概率
print("小尺度概率:", torch.softmax(small_logits, dim=1).detach().numpy())
print("大尺度概率:", torch.softmax(large_logits, dim=1).detach().numpy())
你会发现大尺度logits虽然损失值更小,但可能导致:
- 数值不稳定(可能出现NaN)
- 梯度消失问题(概率过于集中导致梯度变小)
解决方案:适当调整初始化策略或添加批归一化层控制logits尺度。
4.2 类别不平衡时的应对
当某些类别样本极少时,标准交叉熵可能导致模型偏向多数类。解决方案包括:
- 类别加权:
# 假设类别0的样本数是类别1的5倍
weights = torch.tensor([1.0, 5.0])
loss_fn = nn.CrossEntropyLoss(weight=weights)
- 标签平滑(Label Smoothing):
loss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)
# 相当于把真实标签从1.0变为0.9,其余类别从0.0变为0.1/(n_classes-1)
- Focal Loss(对易分类样本降权):
class FocalLoss(nn.Module):
def __init__(self, alpha=1, gamma=2):
super().__init__()
self.alpha = alpha
self.gamma = gamma
def forward(self, inputs, targets):
ce_loss = F.cross_entropy(inputs, targets, reduction='none')
pt = torch.exp(-ce_loss)
loss = self.alpha * (1-pt)**self.gamma * ce_loss
return loss.mean()
4.3 多标签分类的变体
标准交叉熵假设每个样本只属于一个类别。对于多标签问题(如图片中同时包含猫和狗),需要使用:
# 每个标签独立的二分类问题
loss_fn = nn.BCEWithLogitsLoss() # 带sigmoid的二值交叉熵
# 示例
multi_labels = torch.tensor([[1, 0, 1]], dtype=torch.float) # 同时属于类别0和2
multi_logits = torch.randn(1, 3)
loss = loss_fn(multi_logits, multi_labels)
更多推荐



所有评论(0)