为什么你的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)