CE损失函数入门指南:从数学原理到PyTorch实战

1次阅读
没有评论

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

image.webp

从信息论理解交叉熵

交叉熵(Cross-Entropy)本质上是信息论中衡量两个概率分布差异的工具。要理解它,我们需要先了解几个基础概念:

CE 损失函数入门指南:从数学原理到 PyTorch 实战

  1. 信息熵:表示一个事件的不确定性,定义为 (H(p) = -\sum p(x)\log p(x))。例如,硬币正面概率为 0.5 时,熵达到最大值 1。
  2. KL 散度:衡量两个分布 p(真实分布)和 q(预测分布)的差异:(D_{KL}(p||q) = \sum p(x)\log\frac{p(x)}{q(x)})。
  3. 交叉熵:可以拆解为 (H(p,q) = H(p) + D_{KL}(p||q))。由于训练时 p 是固定标签,最小化交叉熵等价于最小化 KL 散度。

为什么分类任务偏爱 CE 而非 MSE

在二分类任务中,假设真实标签 y =1,模型输出概率为 0.01 时:

  • MSE 损失:((1-0.01)^2 = 0.98)
  • CE 损失:(-\log(0.01) \approx 4.6)

CE 对错误预测的惩罚更严厉(梯度更大),尤其当预测概率接近 0 或 1 时。这使得模型能更快修正严重错误的预测。下图展示了两种损失函数的梯度差异:

import numpy as np
import matplotlib.pyplot as plt

x = np.linspace(0.01, 0.99, 100)
mse = (1-x)**2
ce = -np.log(x)

plt.plot(x, mse, label='MSE')
plt.plot(x, ce, label='CE')
plt.legend()
plt.show()

PyTorch/TensorFlow 实战代码

基础用法(带类别权重)

# PyTorch 实现
import torch.nn as nn

# 假设类别 0 和 1 的样本比例为 1:4
weights = torch.tensor([1.0, 0.25])
loss_fn = nn.CrossEntropyLoss(weight=weights)

# 输入格式:batch_size × num_classes
inputs = torch.randn(3, 2)
# 标签格式:batch_size(必须是类别索引)targets = torch.tensor([1, 0, 1])

loss = loss_fn(inputs, targets)
# TensorFlow 实现
import tensorflow as tf

weights = tf.constant([1.0, 0.25])
loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(
    from_logits=True,
    reduction=tf.keras.losses.Reduction.SUM_OVER_BATCH_SIZE
)

# 注意:这里的 weights 是样本级权重,不是类别权重
loss = loss_fn(y_true=[1, 0, 1],
    y_pred=[[0.3, 0.7], [0.8, 0.2], [0.4, 0.6]],
    sample_weight=[0.25, 1.0, 0.25]
)

标签平滑实现

标签平滑(Label Smoothing)通过将硬标签替换为((1-\epsilon)\cdot y + \epsilon/K)(K 为类别数),防止模型对标签过度自信:

# PyTorch 手动实现
epsilon = 0.1
num_classes = 10

# 原始 one-hot 标签
targets = torch.tensor([2, 5])  # 类别索引
smoothed = torch.full((len(targets), num_classes), epsilon/(num_classes-1))
smoothed.scatter_(1, targets.unsqueeze(1), 1-epsilon)

# 使用带 logits 的 CE 损失
loss_fn = nn.CrossEntropyLoss()
logits = model(inputs)
loss = loss_fn(logits, smoothed)

三大避坑指南

  1. 数值稳定性
  2. 永远不要直接将 softmax 输出传入 CE,应先使用 log_softmax 或保持 logits 状态
  3. PyTorch 的 nn.CrossEntropyLoss 已内置优化,等价于nn.LogSoftmax + nn.NLLLoss

  4. 类别不平衡处理

  5. 样本权重(sample_weight)和类别权重(weight)是不同的概念
  6. 当使用采样策略时,需同步调整损失函数的权重参数

  7. 多标签任务适配

  8. 标准 CE 要求单标签分类,多标签需改用nn.BCEWithLogitsLoss
  9. 此时每个类别独立计算 sigmoid 交叉熵

开放性问题

当类别数量极大(如推荐系统中的百万级物品),CE 计算会面临:
– 计算 softmax 分母的归一化项开销巨大
– 负样本梯度被极度稀释

改进方案包括:
– 采样近似(NCE, Negative Sampling)
– 层次 softmax(Hierarchical Softmax)
– 最近提出的噪声对比估计(InfoNCE)

你遇到过哪些 CE 计算的挑战?欢迎分享你的解决方案!

正文完
 0
评论(没有评论)