为什么你的Loss函数不叫“KL散度”?PyTorch分类损失实战指南
为什么你的Loss函数不叫“KL散度”?PyTorch分类损失实战指南
训练分类模型时,我们几乎总是用nn.CrossEntropyLoss()。但翻开信息论教材,衡量分布差异的正统指标是KL散度。既然KL散度才是分布距离的度量,为什么框架的分类损失都叫交叉熵?
真正用到KL散度时(比如知识蒸馏或标签平滑),nn.KLDivLoss的参数设计又容易让人踩坑:input要不要取log?target是什么格式?reduction怎么选?
本文解决三个工程问题:为什么优化交叉熵等于优化KL散度;PyTorch分类Loss如何按任务选型;KL散度在蒸馏和标签平滑中的正确用法。
一、交叉熵与KL散度的优化等价性
测量误差类比
用一把有固定偏差的尺子测零件长度,读数包含零件真实长度(固定)和系统误差(待消除)。交叉熵是带基准的读数,KL散度是扣除基准后的净误差。真实分布P由数据集决定,其熵H(P)是常数,最小化交叉熵即最小化KL散度。
公式表达为 H(P,Q) = H(P) + D_KL(P||Q),H(P)为常数项,优化方向一致。
PyTorch参数映射:
F.kl_div(input=log_Q, target=P)计算 D_KL(P||Q)。input为模型预测的对数概率,target为真实分布的概率值。若target本身已是log-prob(如教师返回log_softmax),应改用F.kl_div(log_Q, log_P.exp())或手写(P * (log_P - log_Q)).sum(),避免对log-prob重复取exp引入数值误差。
定义式验证
不能用F.cross_entropy验证等价性,其内部强制执行log-softmax,传入log-prob会导致双重取log。正确做法是用交叉熵定义式手动计算:
import torch
import torch.nn.functional as F
P = torch.tensor([[0.7, 0.2, 0.1] + [0.0]*7])
Q = torch.rand(1, 10).softmax(dim=-1)
ce_manual = -(P * Q.log()).sum()
kl = F.kl_div(Q.log(), P, reduction='sum')
entropy_P = -(P * P.log()).sum()
print(f"CE - KL = {ce_manual - kl:.6f}, Entropy(P) = {entropy_P:.6f}")
# 输出: CE - KL = 0.892341, Entropy(P) = 0.892341
此代码仅用于验证原理。生产环境中nn.CrossEntropyLoss接收未归一化logits,内部自动做log-softmax保证数值稳定,不要手动对logits取log。
两者梯度也等价,H(P)对模型参数求导为零。交叉熵计算更稳定且无需减常数项,因此成为分类任务的默认选择。
二、PyTorch分类Loss选型
二分类:BCEWithLogitsLoss
输出维度为1或(batch, 1),标签为0/1标量时,优先选BCEWithLogitsLoss。它内部使用log-sum-exp技巧,避免sigmoid极端值溢出;反向传播路径更短,收敛更快;支持pos_weight处理样本不平衡。
# 推荐
criterion = nn.BCEWithLogitsLoss()
loss = criterion(logits, targets)
# 不推荐:BCELoss要求输入为(0,1)概率值,误传logits会导致损失爆炸或归零
criterion = nn.BCELoss()
loss = criterion(torch.sigmoid(logits), targets)
多分类:CrossEntropyLoss
输出维度(batch, C)、C>2、标签为类别索引时,使用CrossEntropyLoss。该函数融合LogSoftmax与NLLLoss,传入未归一化logits即可,不要手动加softmax。
注意两个调参点:ignore_index设为-100可忽略padding等无效标签;weight传入长度为C的张量实现类别级加权。标签必须是LongTensor类型的类别索引,误传one-hot不会报错但结果错误。PyTorch 2.0+支持float型软标签target,但软标签模式与label_smoothing参数互斥,同时使用会报错。自定义软标签需用手动KL或手写交叉熵。
多标签分类:BCEWithLogitsLoss
每个样本可同时属于多个类别时,本质是C个独立二分类问题,使用BCEWithLogitsLoss,标签为(batch, C)的float型multi-hot矩阵。不要用CrossEntropyLoss,其全局softmax会破坏标签独立性。
三、KL散度实战:蒸馏与标签平滑
知识蒸馏
让学生模型输出分布逼近教师软标签分布。
class DistillationLoss(nn.Module):
def __init__(self, temperature=4.0, alpha=0.7):
super().__init__()
self.T = temperature
self.alpha = alpha
# reduction必须指定batchmean,默认mean会使loss随类别数缩放,与T²失配
self.kl = nn.KLDivLoss(reduction='batchmean')
self.ce = nn.CrossEntropyLoss()
def forward(self, student_logits, teacher_logits, hard_labels):
student_log_prob = F.log_softmax(student_logits / self.T, dim=-1)
# 显式detach防止教师计算图残留占用显存
teacher_prob = F.softmax(teacher_logits.detach() / self.T, dim=-1)
# 乘T²补偿高温下的梯度衰减,使soft/hard loss梯度贡献稳定
soft_loss = self.kl(student_log_prob, teacher_prob) * (self.T ** 2)
hard_loss = self.ce(student_logits, hard_labels)
return self.alpha * soft_loss + (1 - self.alpha) * hard_loss
根据Hinton推导,高温极限下蒸馏损失对logits的梯度近似为(1/T²)·(z_s - z_t),乘T²补偿梯度衰减,使soft loss与hard loss梯度贡献稳定。忽略此项会导致soft loss训练中后期失效。alpha通常取0.5~0.9,教师越强、数据越少则越大。
教师模型前向传播需加torch.no_grad()并设为eval模式,代码中已显式detach形成双重保障。
标签平滑
标签平滑将硬目标替换为软目标分布,此时优化交叉熵等价于优化对该软目标的KL散度(差一个常数),它是主损失本身。
均匀平滑可直接用原生API:
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
loss = criterion(logits, hard_labels)
自定义平滑分布时用手动KL实现,需正确处理ignore_index样本,确保其不参与损失和梯度:
def label_smoothing_kl(logits, hard_labels, epsilon=0.1, num_classes=1000):
smooth_targets = torch.full_like(logits, epsilon / num_classes)
valid_mask = (hard_labels >= 0)
safe_labels = hard_labels.clamp(min=0)
smooth_targets.scatter_(1, safe_labels.unsqueeze(1), 1.0 - epsilon + epsilon / num_classes)
log_prob = F.log_softmax(logits, dim=-1)
kl_per_sample = F.kl_div(log_prob, smooth_targets, reduction='none').sum(dim=-1)
# 全ignore batch时返回0,调用方宜检查valid_mask.sum()>0再执行backward
loss = (kl_per_sample * valid_mask).sum() / valid_mask.sum().clamp(min=1)
return loss
训练初期监控smooth_targets熵值,接近0说明平滑未生效,等于log(K)说明退化为均匀分布,合理区间为(0, log(K)),epsilon通常取0.05~0.2。
KL散度使用边界
D_KL(P||Q)中P(x)=0时贡献为0,完全有定义;Q(x)=0且P(x)>0才会数值爆炸。对应nn.KLDivLoss(input=log_Q, target=P),target=0安全,input含log(0)才会出问题。需警惕把模型预测概率Q手动.log()后作为input的写法,生产代码应始终用F.log_softmax。教师输出经softmax后恒正,蒸馏场景极少遇到此问题。
标准分类无软标签时直接用CrossEntropyLoss;需要对称距离度量时选用JS散度或Wasserstein距离。
四、避坑索引与决策矩阵
症状排查索引
| 症状 | 可能原因 | 快速定位 |
|---|---|---|
| Loss NaN/Inf | input含log(0) | 检查是否手动.log()而非F.log_softmax |
| Loss NaN/Inf | sigmoid溢出 | 换用BCEWithLogitsLoss或clamp输入 |
| Loss恒为零/不降 | 温度未补偿 | 蒸馏KL项是否乘T² |
| Loss不降/结果随机 | 标签类型错误 | CrossEntropy target是否为Long索引 |
| Loss不降/梯度异常 | input/target顺序颠倒 | KLDivLoss是否遵循log_Q/P顺序 |
| Loss量级异常 | reduction误用 | KLDivLoss是否指定batchmean |
| 多标签性能崩塌 | 误用CrossEntropy | 多标签应换BCEWithLogitsLoss |
| 泛化差 | 软标签未归一化 | target概率和是否为1 |
| 平滑无效/噪声大 | epsilon不当 | 监控smooth_targets熵值是否在(0,logK) |
完整决策矩阵
| 任务场景 | 推荐Loss | 关键参数 | 输入/标签格式 | 适用场景与备注 | 快速验证方法 |
|---|---|---|---|---|---|
| 二分类 | BCEWithLogitsLoss | pos_weight | logits (B,1) / float 0-1 (B,1) | 标准二分类 | loss∈(0,∞),梯度非零 |
| 多分类 | CrossEntropyLoss | ignore_index, weight, label_smoothing | logits (B,C) / Long索引 (B,) | 软标签与label_smoothing互斥 | 标签近似均匀时随机初始化loss≈log(C);类别不平衡时偏小属正常 |
| 多标签 | BCEWithLogitsLoss | pos_weight | logits (B,C) / float multi-hot (B,C) | 类别不互斥 | 每类独立计算AUC |
| 知识蒸馏 | KLDivLoss + CE | T, alpha, reduction=’batchmean’ | log_student / prob_teacher / hard_labels | 教师logits需detach | soft_loss×T²与hard_loss同量级 |
| 标签平滑(均匀) | CrossEntropyLoss | label_smoothing=ε | 同多分类 | 快速实验、标准分类 | 平滑后验证集acc提升 |
| 标签平滑(自定义) | KLDivLoss | ε, 自定义分布 | log_prob / smooth_prob | 替代CE(label_smoothing),二者不可共用 | smooth_targets熵∈(0,logC) |