ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

交叉熵损失函数:从信息论到PyTorch与GAN实践

交叉熵损失函数:从信息论到PyTorch与GAN实践 大半年里我面试了不少做算法的候选人聊到损失函数时十个人里有八个都能把交叉熵公式默写出来但真问到“为什么分类问题用交叉熵而不是MSE”“GAN的判别器损失里为什么交叉熵没有负号”时能讲清楚的人少之又少。这两个细节恰恰是理解交叉熵的关键也最能看出一个人是在背公式还是真的在用它。如果你想真正弄懂交叉熵而不只是会调一行loss nn.CrossEntropyLoss()这篇文章值得你慢慢看完。我会从信息论的地基讲起一路讲到BCE的设计原理、数值实现、GAN里那个著名的“无负号交叉熵”最后再附上我踩过的一些坑。内容偏原理但也带实操适合正在学深度学习的初学者也适合想补强基础的工程师。1. 从“信息量”到“熵”交叉熵的地基1.1 信息量越不可能的事信息越值钱交叉熵这个名字里有个“熵”字而熵这个概念最早来自热力学后来香农把它引进了信息论用来度量不确定性。不理解信息量就很难真正理解熵和交叉熵所以我先花点篇幅讲这个地基。想象一个场景天气预报说明天晴这基本没什么信息量因为你每天都在经历晴天这件事太“平常”了。但如果预报说明天要下大暴雨这个消息的信息量就很大因为它打破了你的预期带来了更多不确定性。香农把这种直觉量化成了一个公式[ I(x) -\log P(x) ]就是取事件发生的概率再取负对数。概率越小负对数越大信息量越大概率为1的事件负对数为0完全没有信息量。负号是为了保证信息量是正的因为概率在0到1之间直接取对数会得到负数。这个式子虽然简单但它有很实际的意义信息量不是看事件本身重不重要而是看它有多“意外”。一个极端事件的发生比一个普通事件的发生携带了更多信息。机器学习里面的模型训练也是这个道理——一个低概率样本出现时模型会“意外”这时候更新参数的幅度和方向恰好就需要这种意外感来驱动。1.2 熵对信息量求期望有了单个事件的信息量如果我们想描述一个整个概率分布的平均信息量就需要把每个可能事件的信息量按出现的概率加权求和这个加权平均值就是熵。[ H(P) -\sum_{x} P(x) \log P(x) ]熵衡量的是一个分布本身的不确定性。向一个均匀分布的骰子每个面概率都是1/6信息量是均匀的不确定性最大向一枚被做了手脚的硬币如果正面概率是0.99那么每抛一次几乎都能猜到结果不确定性很小熵就低。在机器学习里熵经常被当作“混乱程度”的代名词。模型预测出的概率分布越接近均匀分布熵越高说明模型对答案越不确定越接近one-hot分布熵越低说明模型越有把握。注意熵的最小值不是恒定的0只有某个事件概率为1而其他事件概率为0时熵才是0。对于C分类问题熵的最小值就是0最大值是log C。1.3 KL散度用交叉熵衡量两个分布的距离现在我们进入核心问题模型预测的分布Q距离真实分布P有多远KL散度就是用来做这件事的。[ D_{KL}(P \parallel Q) \sum_{x} P(x) \log \frac{P(x)}{Q(x)} ]把除法拆开就变成[ D_{KL}(P \parallel Q) \sum_{x} P(x) \log P(x) - \sum_{x} P(x) \log Q(x) ]第一项是真实分布P的熵的相反数它是固定不变的因为P是我们手里的标签分布不会随模型训练改变。第二项也就是 (-\sum P(x) \log Q(x))就是大名鼎鼎的交叉熵。所以“让预测分布Q尽量接近真实分布P”这件事等价于“最小化KL散度”而由于KL散度中的第一项是常数就又等价于“最小化交叉熵”。这就是为什么深度学习里大家都在最小化交叉熵——它没有直接去计算分布间的距离而是在最小化和距离等价的那一部分目标。这里需要特别注意KL散度并不是真正的“距离”它不对称。(D_{KL}(P \parallel Q)) 不等于 (D_{KL}(Q \parallel P))。前者是用P做加权关注P中概率大的区域拟合得好不好后者是用Q做加权关注Q中概率大的区域。在训练生成模型和蒸馏模型时这个不对称性经常带来让人头疼的问题也正因如此后来才出现了各种对称化的变体。2. 交叉熵损失的设计原理为什么分类任务都爱它2.1 从极大似然出发交叉熵就是负对数似然深度学习里的交叉熵损失从来不是凭空发明的它是从极大似然估计推出来的。用一句话说我们想让模型在所有训练样本上产生正确标签的概率乘积最大。给定N个样本模型预测出每个样本属于正确类别c的概率是 (P(c|x_i))那么所有这些概率的乘积就是似然函数[ L \prod_{i1}^{N} P(c_i|x_i) ]连乘计算容易数值下溢取对数变成连加再取负号变成最小化目标[ Loss -\frac{1}{N} \sum_{i1}^{N} \log P(c_i|x_i) ]这个过程就是负对数似然NLL。现在对照一下交叉熵公式当真实分布P是one-hot向量时(P(c)1)其他位置为0交叉熵的求和只剩下一项就是 (-\log Q(c))。这说明在分类任务里交叉熵和负对数似然就是同一个东西。这也解释了为什么交叉熵损失在这种情况下被吐槽为“只看正确类别的预测概率不管错误类别的分布”。当某个样本的真实类别概率很低时梯度会很大模型会被快速拉回正确方向当真实类别概率接近1时梯度变小模型收敛得很平滑。2.2 对比MSE为什么回归损失不适合分类很多人刚接触深度学习时会有个疑问MSE也能度量预测和标签的差距为什么分类任务不用它我回答这个问题从来不说“经验上不好用”而是从函数曲线上直接看。假设二分类问题标签y是0或1模型输出经过sigmoid后得到预测概率p。如果用MSE[ Loss (p - y)^2 ]梯度是 (2(p-y) \cdot p(1-p))。问题出在sigmoid的导数 (p(1-p)) 上。当预测概率p接近0或1时sigmoid曲线趋于平缓梯度趋近于0这本来不是什么大问题因为此时预测已经接近正确标签了确实应该慢下来。但麻烦的是在p接近0.5的时候模型还没分对梯度却已经很小了。这会让训练在初始阶段慢得像蜗牛爬。交叉熵没有这个问题。代入二分类交叉熵BCE的梯度公式可以发现梯度正比于 ((p-y))不受sigmoid导数影响。这就是为什么分类任务里交叉熵的训练速度和稳定性远好于MSE。如果一定要做个直观对比可以看这张表对比维度交叉熵MSE错误区域梯度大快速纠正小收敛慢与softmax/sigmoid组合梯度形式简单数值稳定容易饱和梯度消失对离群样本的容忍度对错误概率很敏感对离群误差惩罚大典型适用场景分类、多标签、生成对抗训练回归、自监督重构所以当你再看到某个新手的代码里分类用了MSE不要只觉得是精度问题本质上是训练动力学的问题。2.3 BCE的设计原理二分类的交叉熵长什么样二分类是分类问题的基础情况它的交叉熵形式就是BCEBinary Cross Entropy。先回顾一般交叉熵真实分布P是one-hot时交叉熵只取正确类别的负对数概率。二分类里我们可以把两个类别看成一个概率分布(P(y1)y)(P(y0)1-y)。模型预测的分布为 (Q(y1)p)(Q(y0)1-p)。把它写进交叉熵公式[ H(P,Q) -[y \log p (1-y)\log(1-p)] ]这就是BCE。它的设计原理非常清晰每个样本只包含两个事件这两个事件互补。模型输出p作为正例概率1-p自然是负例概率。真实标签y如果是1那么这一项变成 (-\log p)模型就应该尽量提高py是0就变成 (-\log(1-p))模型就应该尽量压低p。这里有一个我经常在代码里看到的错误把BCE理解成“两个类别的交叉熵之和”于是写了-(y log p (1-y) log(1-p))这个公式没错但要注意p本身必须是sigmoid输出或者直接换用带logits的BCEWithLogitsLoss而不是把模型原始的线性输出直接塞进去算。BCE从设计上就假设了输出是概率值所以配合sigmoid使用是标准操作。现在的框架大都提供BCEWithLogitsLoss内部把sigmoid和BCE合并在一起算好处一是数值上更稳定二是计算更快这个后文还会展开。3. 实操落地PyTorch里怎么用最稳3.1 softmax和交叉熵是天作之合合并计算的数值稳定多分类问题的标准套路是最后一层接softmax把logits转换成概率然后和交叉熵一起算。但如果你直接在代码里写两步先softmax再log再相乘数值上很容易出问题。比如某个logit特别大比如20softmax后这个类别的概率可能接近1再取log接近0看起来没毛病。但极端情况下logits相对差距很大softmax内部的指数函数会直接溢出或者出现inf然后log(0)得到-inf。PyTorch的CrossEntropyLoss对此做了合并处理它直接接受logits作为输入内部把softmax和log合起来算用log-sum-exp的技巧保证数值稳定而不溢出。为什么能合起来算因为[ -\log\left(\frac{e^{z_c}}{\sum_j e^{z_j}}\right) -z_c \log \sum_j e^{z_j} ]把softmax的除法取对数的过程简化成了减法和logsumexp这样即使某个z_j非常大logsumexp也能稳定计算。所以你在写代码时如果用的是CrossEntropyLoss千万不要在模型输出后面自己再套一层softmax否则你等于做了两次softmax逻辑错误不说数值上还白白引入了一次不必要的计算。3.2 三种常用API的区别CrossEntropyLoss vs NLLLoss vs BCEWithLogitsLoss我见过不少初学者把这些API搞混这里直接列清楚API输入要求适用场景nn.CrossEntropyLoss模型输出原始logits标签传类别索引单标签多分类nn.NLLLoss模型输出已经经过log_softmax的结果单标签多分类传统写法nn.BCEWithLogitsLoss模型输出原始logits标签传0/1浮点数二分类、多标签分类nn.BCELoss模型输出概率值已过sigmoid二分类、多标签不推荐直接裸用NLLLoss其实才是“负对数似然损失”的本体CrossEntropyLoss在PyTorch里的实现就是log_softmax NLLLoss的封装。如果你看到老代码用nn.LogSoftmax()之后接nn.NLLLoss()不要觉得奇怪那只是同一个东西的另一种写法。写一段代码示例展示怎么在训练里正确使用import torch import torch.nn as nn import torch.nn.functional as F # 单标签多分类直接用CrossEntropyLoss model nn.Linear(64, 10) logits model(x) # shape: (batch, 10) loss_fn nn.CrossEntropyLoss() loss loss_fn(logits, target) # target shape: (batch,)元素是0~9的索引 # 二分类BCEWithLogitsLoss binary_model nn.Linear(64, 1) logits binary_model(x) # shape: (batch, 1) loss_fn nn.BCEWithLogitsLoss() loss loss_fn(logits, target.float()) # target: (batch, 1)元素0或1 # 多标签分类同一个BCEWithLogitsLosstarget是one-hot或0/1向量 multilabel_logits model(x) # shape: (batch, num_labels) loss_fn nn.BCEWithLogitsLoss() loss loss_fn(multilabel_logits, multilabel_target.float()) # multilabel_target同shape注意target的类型和形状。CrossEntropyLoss要的是LongTensor的类别索引不是one-hotBCEWithLogitsLoss要的是FloatTensor而且形状要和logits完全一致。这两点写错是最常见的运行时报错来源。3.3 类别不平衡、标签平滑等工程调整真实数据集很少是均匀分布的。有时候正样本只占5%负样本占95%直接拿BCE去算模型会学成一个“什么都预测成负类”的分类器因为这样能轻易把loss降到很低。处理类别不平衡有几个常用技巧按优先级排列调整类别权重BCEWithLogitsLoss(pos_weight...)给正样本更高的权重。pos_weight可以设为负数样本数除以正样本数。这个做法在信息检索和推荐场景里实测最直接有效。对少数类做上采样简单粗暴但容易过拟合配合数据增强用会好很多。换用Focal Loss在交叉熵前面加一个调制因子 ((1-p)^\gamma)让模型把注意力集中在难分样本上。Focal Loss本质上仍是交叉熵只是对每个样本做了重新加权这个权重和模型当前对样本的置信度有关。标签平滑是另一个我在分类模型里必开的小工具。它的原理很简单把one-hot标签从1和0变成 (1-\epsilon) 和 (\epsilon/(C-1))。比如10分类(\epsilon0.1)那么正确类别的标签变成0.9其他9个类别各分到约0.011。这样做的直接效果是模型不会在某个类别上无限增大logits差值从而提升泛化性和校准性。标准的CrossEntropyLoss本身没有内置label smoothing参数PyTorch从1.10开始为它加了这个参数用法是loss_fn nn.CrossEntropyLoss(label_smoothing0.1) loss loss_fn(logits, target)如果模型已经训练好了想在推理时得到更靠谱的置信度也可以考虑Temperature Scaling这个属于后处理校准不细说但记住一点交叉熵训练出来的模型预测概率不一定等于真实概率这不奇怪。4. 原始GAN公式的交叉熵为什么没有负号4.1 从GAN的目标函数看D的视角如果你看过原始GAN论文一定对下面这个目标函数有印象[ \min_G \max_D V(D,G) \mathbb{E}{x\sim P{data}}[\log D(x)] \mathbb{E}_{z\sim P_z}[\log(1 - D(G(z)))] ]很多人第一次接触时都懵了交叉熵不是在log前面有个负号吗为什么这里的log D(x)没有负号其实这里藏着GAN最大的一个设计巧思也藏着最容易被初学者误解的点。先只从判别器D的角度看。D的任务是分辨真实样本和生成样本。真实样本的标签是1生成样本的标签是0。如果我们给D套一个标准的BCE损失那它应该最小化[ L_D -\mathbb{E}{x\sim P{data}}[\log D(x)] - \mathbb{E}_{z\sim P_z}[\log(1 - D(G(z)))] ]这个式子里是有负号的。但GAN的作者写的目标函数里没有负号因为GAN的整个框架不采用“单独给D算BCE loss再反向传播”的写法而是写成了一个极大极小博弈。max_D V(D,G)意思是D要最大化V而V里恰好就是标准的BCE公式去掉负号后的样子。D最大化这个V等价于最小化加负号的BCE。所以不是没有负号而是负号被“极大化”这个动作吸收掉了。4.2 没有负号那D是怎么更新的实际用PyTorch训练GAN时D的更新代码通常是这样的# 真实样本的loss real_loss -torch.mean(torch.log(discriminator(real_data))) # 这里加负号 # 生成样本的loss fake_loss -torch.mean(torch.log(1 - discriminator(fake_data))) # 总loss d_loss real_loss fake_loss或者更常见的是直接调用BCEWithLogitsLoss给真实样本标签1、给生成样本标签0d_loss bce_loss(discriminator(real_data), torch.ones_like(...)) \ bce_loss(discriminator(fake_data), torch.zeros_like(...))你看真正写代码的时候还是要加负号因为优化器只认“下降方向”它默认做的是minimize。论文里写max_D是数学语言代码里落地的时候你永远要把它换成minimize某个等价的loss。如果谁真的照着论文式子不加负号去训练D会朝着反方向更新模型根本学不动。不过这里有一个值得注意的细节论文的V里第一项是真实样本的期望第二项是生成样本的期望。D想尽量区分真假所以最大化第一项让D(real)接近1和最大化第二项的log(1-D(G(z)))让D(fake)接近0。而生成器G想骗过D所以它希望最小化第二项也就是让D(G(z))接近1。G不能控制真实样本那一项那一项跟G无关。4.3 优化器把负号藏在哪minimize和maximize的等价关系我们来把这件事彻底说透。假设D的standard BCE损失是 (L_D)则[ L_D -\mathbb{E}[\log D(x)] - \mathbb{E}[\log(1 - D(G(z)))] ]而论文的V是[ V \mathbb{E}[\log D(x)] \mathbb{E}[\log(1 - D(G(z)))] ]显然 (V -L_D)。最大化V就是最小化 (L_D)。这就是为什么论文里没有负号本质上是把“最小化损失”变成了“最大化收益”来写这在博弈论的框架里叫收益函数而不是损失函数。还有一个常见问题很多人会用这个替代写法训练G最小化 (-\log D(G(z)))而不是原始论文里的 (\log(1-D(G(z))))。原因很简单原始写法在D太强时梯度几乎为0G学不动也就是常说的“饱和”而-log D(G(z))在D(G(z))接近0时梯度非常大能给G提供更强的信号。这叫“非饱和损失”是Goodfellow在同一篇论文里就提出来的改进建议只是大家默认不提为论文的修改。实操上我建议GAN的判别器一律用BCEWithLogitsLoss生成器要么用-log(D(G(z)))的非饱和形式要么用最小二乘等更稳的损失。到了现代GAN比如WGAN、WGAN-GP、StyleGAN时代已经很少有人直接用原始交叉熵形式的损失了因为训练不稳定问题太突出。但理解原始GAN里负号的来龙去脉能帮你把“损失函数在数学表达和代码实现之间如何转换”这个底层能力彻底打通。这里可以用一个很小的代码片段来印证这种等价关系# 判别器论文写法maximize V(D,G) # 优化器要做的是minimize所以要写 -V d_loss -torch.mean(torch.log(D_real)) - torch.mean(torch.log(1 - D_fake)) # 上面这行等价于 d_loss torch.mean(F.binary_cross_entropy(D_real, torch.ones_like(D_real))) \ torch.mean(F.binary_cross_entropy(D_fake, torch.zeros_like(D_fake)))两个写法在数学上完全一致只是前者的数值稳定性差因为log(0)会出问题所以实际工程里一定用后者或者用BCEWithLogitsLoss。5. 常见问题与排查技巧实录5.1 Loss变成负数或NaN交叉熵的正常取值范围是大于等于0如果出现负数几乎可以确定是输入有问题。我排查loss异常时习惯按下面顺序检查标签是不是从0开始的连续整数如果标签里混进了-1或大于类别数的值CrossEntropyLoss的index会越界轻则报错重则静默算出错值。模型输出有没有可能包含NaN尤其是用了自定义网络时检查最后一层有没有归一化、有没有除零风险。BatchNorm在batch size过小时也会出数值问题。学习率是不是太大了一开始就出现NaN多半是学习率炸了试着把学习率除以10再看。数据里有没有inf一颗像素值出现inf就能让整个batch的loss变成NaN。如果loss是NaN别急着改loss函数90%的情况是上游输入或者数值稳定性出了问题而不是交叉熵本身的问题。5.2 One-hot编码和索引标签搞混nn.CrossEntropyLoss接收的target是[0, C-1]的整数索引而不是one-hot编码。很多人习惯把标签用F.one_hot转成向量然后直接喂给CrossEntropyLoss结果报错。如果你想用one-hot标签可以手动实现交叉熵-torch.sum(target_onehot * F.log_softmax(logits, dim-1), dim-1)或者在PyTorch里直接传索引让框架内部去处理。BCEWithLogitsLoss就反过来它要求target是0/1的浮点张量形状和logits完全一致。二分类里常见的一个坑是模型输出形状是(batch,)标签形状是(batch, 1)直接算loss广播规则会把维度搞错数值不会报错但结果完全不对。解决方法是统一label.view(-1, 1)。5.3 多标签和多分类被当成一回事多分类Multi-class和多标签Multi-label是两个完全不同的任务对应不同的损失函数。多分类假设每张图片属于且仅属于一个类别输出层用softmax损失用CrossEntropyLoss标签是索引。多标签假设每张图片可以同时拥有多个属性比如“蓝天、沙滩、人物”三个标签可以同时为1输出层每个位置用sigmoid损失用BCEWithLogitsLoss标签是一个0/1向量可能全0也可能全1。如果把多标签任务误用CrossEntropyLoss模型会强行在标签之间做竞争导致每个样本只能预测出一个标签损失函数和任务需求直接错位。我排查这类问题时只要看一眼训练代码用的是哪个loss就能判断项目方案是不是一开始就选错了框架。5.4 观察loss曲线时我还喜欢顺带看这几样东西交叉熵loss本身能反映的信息有限尤其是当准确率已经很高时loss的绝对值很难直观反映模型好坏。我训练时通常会额外记录这些量平均置信度对正确类别预测概率的均值。如果这个值持续上升说明模型在朝正确的方向走。熵对所有类别的预测概率求熵。熵高说明模型犹豫不决训练后期如果熵迟迟降不下来说明类别区分度不够。logits的分布直接看logits的均值和方差。Logits过大或过小往往意味着标签平滑没开或者初始化不合理。有一次我训练一个二分类模型准确率卡在90%上不去loss曲线看上去也挺正常。后来我打印了中间层特征的分布发现模型在特征空间里已经能基本分开两个类但最后sigmoid之前的logits偏差很大正样本的logits平均比负样本高出一大截。这其实是类别不平衡导致的logit偏移我用pos_weight调整BCE的权重后准确率才重新动了起来。所以交叉熵不仅仅是公式里的那一行它和标签分布、模型输出分布、数值精度都是强相关的。把这个损失函数当成一个反馈循环去看比单独调一个loss要有效得多。我个人在实际使用中最深的体会是交叉熵的数学形式虽然简单但它连接了信息论、概率论和优化理论是理解深度学习模型训练的绝佳切入点。每当我看到一个新的网络结构第一反应就是去推它的损失函数看它到底在优化什么分布、在哪个维度上约束模型。如果你也能做到这一点那以后看任何论文、写任何模型思路都会清晰很多。最后再分享一个小习惯写完训练代码后先打印一个batch的loss做数值校验拿一个很小的玩具样本集过拟合一遍确认loss确实在下降再上全量数据能省下很多排查时间。
返回列表