从代码实践反推交叉熵: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的实现有几个关键特性常被忽略:

  1. 内置softmax:实际上不需要在模型最后一层添加softmax
  2. 输入格式:接受原始logits(未归一化的分数)而非概率
  3. 标签处理:可以使用类索引(如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 类别不平衡时的应对

当某些类别样本极少时,标准交叉熵可能导致模型偏向多数类。解决方案包括:

  1. 类别加权
# 假设类别0的样本数是类别1的5倍
weights = torch.tensor([1.0, 5.0])
loss_fn = nn.CrossEntropyLoss(weight=weights)
  1. 标签平滑(Label Smoothing):
loss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)
# 相当于把真实标签从1.0变为0.9,其余类别从0.0变为0.1/(n_classes-1)
  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)
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐