共计 1380 个字符,预计需要花费 4 分钟才能阅读完成。
背景与数值稳定性问题
在 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$),会出现以下问题:
- 指数溢出 :
exp(z_i)超过 float32 表示范围(3.4e38) - 梯度异常:根据链式法则,梯度计算包含 $p_i(1-p_i)$ 项,极端情况下会导致梯度消失
PyTorch 实现对比
原生 nn.CrossEntropyLoss
PyTorch 内部采用以下优化策略:
- log-sum-exp 技巧:
$$\log\sum e^{z_i} = \max(z) + \log\sum e^{z_i – \max(z)}$$ - 自动处理 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 数据集上对比:
- 损失值对比实验
| Logits 范围 | 原生 CE Loss | 手动稳定实现 |
|---|---|---|
| [-100,100] | 2.302 | 2.302 |
| [1e3,1e4] | NaN | 8.214 |
- 梯度幅度对比

工程实践建议
- 批处理优化
- 使用
torch.bmm加速矩阵运算 -
避免在损失计算中创建临时张量
-
混合精度训练
with torch.cuda.amp.autocast(): # 需要强制 float32 的运算 loss = stable_ce_loss(logits.float(), labels) -
分布式训练
- 使用
all_reduce同步梯度 - 注意各卡 logits 的独立归一化
延伸思考
- 为什么 PyTorch 的 CrossEntropyLoss 默认不进行 logits 截断?
- 在多标签分类任务中如何修改交叉熵实现?
- 当类别数量扩展到 1000 类时,需要哪些额外优化?
参考文献
- PyTorch 官方文档 – CrossEntropyLoss 实现
- 《Deep Learning》Chapter 4.1
- IEEE 754 浮点数标准
正文完
发表至: 未分类
近一天内
