大半年里我面试了不少做算法的候选人,聊到损失函数时,十个人里有八个都能把交叉熵公式默写出来,但真问到“为什么分类问题用交叉熵而不是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_{i=1}^{N} P(c_i|x_i) ]
连乘计算容易数值下溢,取对数变成连加,再取负号变成最小化目标:
[ Loss = -\frac{1}{N} \sum_{i=1}^{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的设计原理:二分类的交叉熵长什么样
二分类是分类问题的基础情况,它的交叉熵形式就是BCE(Binary Cross Entropy)。
先回顾一般交叉熵:真实分布P是one-hot时,交叉熵只取正确类别的负对数概率。二分类里,我们可以把两个类别看成一个概率分布:(P(y=1)=y),(P(y=0)=1-y)。模型预测的分布为 (Q(y=1)=p),(Q(y=0)=1-p)。把它写进交叉熵公式:
[ H(P,Q) = -[y \log p + (1-y)\log(1-p)] ]
这就是BCE。它的设计原理非常清晰:每个样本只包含两个事件,这两个事件互补。模型输出p作为正例概率,1-p自然是负例概率。真实标签y如果是1,那么这一项变成 (-\log p),模型就应该尽量提高p;y是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特别大,比如20,softmax后这个类别的概率可能接近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 # 多标签分类:同一个BCEWithLogitsLoss,target是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-hot;BCEWithLogitsLoss要的是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分类,(\epsilon=0.1),那么正确类别的标签变成0.9,其他9个类别各分到约0.011。这样做的直接效果是模型不会在某个类别上无限增大logits差值,从而提升泛化性和校准性。
标准的CrossEntropyLoss本身没有内置label smoothing参数,PyTorch从1.10开始为它加了这个参数,用法是:
loss_fn = nn.CrossEntropyLoss(label_smoothing=0.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、给生成样本标签0:
d_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太强时梯度几乎为0,G学不动,也就是常说的“饱和”;而-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确实在下降再上全量数据,能省下很多排查时间。