10分类交叉熵损失函数:原理剖析与PyTorch实战避坑指南

1次阅读
没有评论

共计 1380 个字符,预计需要花费 4 分钟才能阅读完成。

image.webp

背景与数值稳定性问题

在 10 分类任务中,交叉熵损失函数定义为:

$$\mathcal{L} = -\sum_{i=1}^{10} y_i \log(p_i)$$

其中 $p_i$ 通过 softmax 计算得到:

$$p_i = \frac{e^{z_i}}{\sum_{j=1}^{10} e^{z_j}}$$

当 logits 值 $z_i$ 过大时(如 >$10^3$),会出现以下问题:

  1. 指数溢出 exp(z_i) 超过 float32 表示范围(3.4e38)
  2. 梯度异常:根据链式法则,梯度计算包含 $p_i(1-p_i)$ 项,极端情况下会导致梯度消失

PyTorch 实现对比

原生 nn.CrossEntropyLoss

PyTorch 内部采用以下优化策略:

  1. log-sum-exp 技巧:
    $$\log\sum e^{z_i} = \max(z) + \log\sum e^{z_i – \max(z)}$$
  2. 自动处理 logits 维度

手动实现常见问题

# 危险实现(无数值稳定处理)def unsafe_ce_loss(logits, labels):
    probs = torch.softmax(logits, dim=-1)  # 可能溢出
    return -torch.log(probs.gather(1, labels))

数值稳定实现方案

def stable_ce_loss(logits, labels, eps=1e-8):
    """
    logits: [batch_size, num_classes]
    labels: [batch_size]
    数学依据:log(softmax(x)) = x - log(sum(exp(x)))
    """
    # 1. 数值归一化(关键步骤)logits = logits - torch.max(logits, dim=1, keepdim=True)[0]

    # 2. 稳定计算 log_softmax
    exp_logits = torch.exp(logits)
    log_probs = logits - torch.log(exp_logits.sum(dim=1, keepdim=True) + eps)

    # 3. 防御性维度检查
    if labels.dim() == 1:
        labels = labels.unsqueeze(1)

    # 4. 交叉熵计算
    nll_loss = -log_probs.gather(1, labels)
    return nll_loss.mean()

实验验证

在 MNIST-10 数据集上对比:

  1. 损失值对比实验
Logits 范围 原生 CE Loss 手动稳定实现
[-100,100] 2.302 2.302
[1e3,1e4] NaN 8.214
  1. 梯度幅度对比

10 分类交叉熵损失函数:原理剖析与 PyTorch 实战避坑指南

工程实践建议

  1. 批处理优化
  2. 使用 torch.bmm 加速矩阵运算
  3. 避免在损失计算中创建临时张量

  4. 混合精度训练

    with torch.cuda.amp.autocast():
        # 需要强制 float32 的运算
        loss = stable_ce_loss(logits.float(), labels)

  5. 分布式训练

  6. 使用 all_reduce 同步梯度
  7. 注意各卡 logits 的独立归一化

延伸思考

  1. 为什么 PyTorch 的 CrossEntropyLoss 默认不进行 logits 截断?
  2. 在多标签分类任务中如何修改交叉熵实现?
  3. 当类别数量扩展到 1000 类时,需要哪些额外优化?

参考文献

  1. PyTorch 官方文档 – CrossEntropyLoss 实现
  2. 《Deep Learning》Chapter 4.1
  3. IEEE 754 浮点数标准
正文完
 0
评论(没有评论)