交叉熵、信息熵与交叉熵损失:从原理到PyTorch实现
发布时间:2026/9/24 20:31:19 作者:尧图编辑部 阅读量:1,286

做AI项目的人八成都有过这种体验模型代码里一行nn.CrossEntropyLoss()用得贼溜可一旦被人追问“交叉熵到底是什么意思为什么分类任务一定要用它”立刻就有点发虚。我当年也是这么过来的熵、交叉熵、交叉熵损失这三个词听着像亲兄弟脑子里的印象却跟陌生人一样。这篇是大白话AI答疑的第7篇咱们不搞数学论文就用聊天的方式把三兄弟的来历、关系、计算方式彻底捋一遍最后我还会带几组手算示例和PyTorch验证保证你看完能自己动手算出来。这篇内容适合三类人看一是刚入门的AI初学者被公式劝退但还想搞明白原理二是会调包但不敢下钻细节的工程师想补上这块短板三是准备面试的人交叉熵几乎属于必考题搞懂之后至少能聊十分钟不冷场。我尽量少用“显然”“易得”这种让人翻白眼的词遇到公式就拆开讲人话。1. 先从信息量说起熵到底在衡量什么1.1 越“意外”的消息信息量越大想理解熵先要理解单个事件的信息量。举个例子大晴天有人告诉你“今天太阳从东边出来了”你会觉得这是一句废话因为它带给你的信息量约等于零。但要是有人说“明天彩票中奖号码是你的生日”你的反应绝对是被震惊到这句话的信息量直接拉满。信息量的大小本质上取决于一件事发生的概率。概率越低发生时带来的惊讶程度越高信息量就越大。香农把这个直觉写成了公式I(x) -log(p(x))p(x)是事件发生的概率。取负号是因为概率小于等于1对数本身是负的负负得正让信息量是个正数。为什么一定要用对数因为对数有个很好的特性两个独立事件同时发生的总信息量等于各自信息量相加。比如“明天晴天”和“股票涨停”如果互相独立两件事同时发生的概率是各自概率相乘取对数之后就变成了相加。这样信息量就拥有了“可加性”处理起来非常方便。底数用2时信息量单位是bit经典例子是抛硬币。公平硬币正面概率0.5那么出现正面的信息量就是-log2(0.5)1 bit。这代表你需要1位二进制数就能编码这个结果。如果某件事概率是0.25信息量就是2 bit需要2位二进制数才能编码。也就是说信息量越大的事件在通信编码时需要的bit数就越多。1.2 信息熵所有可能结果的平均信息量现实中一个随机事件的输出往往不止两种可能。今天的天气可能是晴、雨、雪模型预测一张图片可能属于猫、狗、鸟。每个可能结果都有自己的概率也都有自己的信息量。信息熵就是把所有可能结果的信息量按概率加权平均得到整个随机系统“平均需要多少信息才能描述清楚”。公式长这样H(P) -Σ p_i * log(p_i)这个公式看着唬人翻译成人话就是把每个结果的概率p_i乘以它的信息量-log(p_i)全部加起来。它度量的是一个概率分布本身的不确定性大小。我拿三个案例给你感受一下场景概率分布熵用ln计算单位nat含义确定性事件[1, 0]0毫无不确定性平均信息量是0公平硬币[0.5, 0.5]0.693不确定性最大完全无法预测偏硬币[0.9, 0.1]0.325虽然有意外但整体比较确定计算过程很简单公平硬币就是-0.5*ln0.5-0.5*ln0.5 0.693。偏一点的硬币-0.9*ln0.9-0.1*ln0.1 0.325。你会发现分布越均匀熵越大分布越极端熵越小。如果某个结果概率变成了1其他全是0熵就是0代表没有任何信息需要传递。从这个角度理解熵就是“我们对一个随机系统有多无知”的量化表达。1.3 从编码角度看熵这是交叉熵的地基熵还有一个特别实用的解释它是无损编码一个随机变量所需的最短平均编码长度下界。假设我们要把天气结果编码成二进制消息发给远方的人。如果晴天和雨天概率各占0.5那你必须用1位二进制数区分两种状态比如0代表晴天、1代表雨天平均编码长度就是1 bit。但如果晴天概率是0.9、雨天是0.1你有没有更省的办法有给晴天编码0给雨天编码1平均长度仍然是1 bit。更好的做法是用霍夫曼编码或香农编码让高概率事件用短码晴天用01位雨天用11位并不能继续压缩因为只有两个符号。如果有四个符号概率分别为0.5、0.25、0.125、0.125设计编码0、10、110、111平均长度是0.5*10.25*20.125*30.125*31.75 bit恰好等于该分布的熵。无论你怎么设计编码都不可能低于1.75 bit这就是熵的意义——它是信息压缩的极限。这个“编码”思想非常重要。交叉熵之所以叫“交叉”本质上就是“我用一套概率分布去编码另一套概率分布时需要付出多少代价”。后面马上就会用上。2. 交叉熵两个分布之间的“代沟”有多大2.1 从熵到交叉熵你拿着错误的码本现在我们把场景升级。假设真实世界里的天气概率分布是P但我们不知道手里只有一套自己估计的概率分布Q。如果直接用Q设计出来的编码规则去描述P产生的数据平均每个消息需要的比特数就是交叉熵。公式H(P, Q) -Σ p_i * log(q_i)你注意区别信息熵里对数值后面跟的是p_i交叉熵里对数值后面跟的是q_i。也就是说概率来自P但编码长度是用Q预估的。用生活化类比P是真实语言Q是一本翻译错误百出的字典你用这本烂字典给一篇文章编码结果长度一定比用原著语言编码要长。一个重要不等式是交叉熵永远大于等于信息熵。原因很简单任何系统的熵都是它的理论最短平均编码长度你用一个不准确的分布去编码只会让编码长度变长不可能更短。只有Q和P完全一致时交叉熵才等于熵。2.2 KL散度交叉熵和熵的差值既然交叉熵一定大于等于熵那这两者的差值就非常值得研究。这个差值就是KL散度也叫相对熵KL(P||Q) H(P, Q) - H(P) Σ p_i * log(p_i / q_i)KL散度衡量的是“用Q去近似P时额外浪费的信息量”。如果Q和P完全一样KL散度为0差异越大KL越大。它满足非负性但它不是距离因为它不对称KL(P||Q)不等于KL(Q||P)。所以严格说它叫“散度”而不是“距离”。比如一个真实的公平硬币P[0.5, 0.5]你的估计Q[0.9, 0.1]计算KL(P||Q)0.5*log(0.5/0.9)0.5*log(0.5/0.1)0.5*ln0.55560.5*ln50.5*(-0.5878)0.5*(1.6094)0.5108。意味着你的错误假设导致平均每个消息多花0.5108 nat的信息。如果我们用Q直接算交叉熵H(P,Q)0.5*ln(1/0.9)0.5*ln(1/0.1)0.05271.15131.204而P本身熵是0.693差值刚好是KL的0.5108。这个闭环你就理解得很扎实了。2.3 模型训练时为什么总提交叉熵而不是KL散度做监督学习时真实标签分布P是固定的模型输出分布Q是不断变化的。我们想让Q尽量接近P也就是最小化KL散度。但仔细看公式KL(P||Q) H(P, Q) - H(P)H(P)是真实分布的熵它完全取决于数据本身与模型参数无关。在一个训练迭代里它就是个常数。所以最小化KL散度和最小化交叉熵在参数更新层面是等价问题。既然如此为什么所有人都在用交叉熵而不直接写KL散度一是交叉熵表达式更简洁二是在求梯度时少算一项常数不影响方向还能省点计算量。更重要的是在分类任务里真实标签通常是one-hot向量此时H(P)等于0KL散度就直接退化成了交叉熵两个数完全一样。所以你可以把交叉熵损失理解成“衡量预测分布和真实标签分布之间代沟”的工具。3. 交叉熵损失分类模型里的那个loss到底是什么3.1 单分类case损失就是正确类别的负对数概率终于说到模型训练里的交叉熵损失了。分类任务里每个训练样本的标签通常用one-hot向量表示比如三分类任务中类别0的标签是[1,0,0]类别1是[0,1,0]类别2是[0,0,1]。在这个前提下交叉熵损失公式会迎来一次剧变L -Σ y_i * log(q_i)由于one-hot向量只有一个位置是1其他位置全是0相乘之后只剩下一个非零项L -log(q_c)这里的q_c是模型给那个正确类别的预测概率。换句话说每个样本的交叉熵损失就是“正确类别预测概率”的负对数。如果模型对正确类别预测得越准比如q_c0.95它的负对数-log0.95≈0.0513损失很小如果模型完全没把握q_c0.1损失就是-log0.1≈2.3026损失很大惩罚非常激烈。这也是为什么你在训练日志里看到的loss下降本质上是模型在不断逼自己把正确类别的预测概率往1靠近。对单分类问题来说交叉熵损失和最大似然估计里的负对数似然是对得上的所以它有一个很硬的理论基础。3.2 多分类和softmax手拉手模型最后一层输出的是一组logits也就是没有归一化的得分向量比如[2.0, 1.0, 0.1]。这些值可以是负数也可以是大于1的数直接放到交叉熵里没有概率意义。所以要先通过softmax把它变成一组加起来等于1的概率q_i exp(z_i) / Σ exp(z_j)softmax做的事就是“让大的更大、小的更小”同时保证所有输出在(0,1)区间。举个例子logits[2.0, 1.0, 0.1]exp(2.0)7.389exp(1.0)2.718exp(0.1)1.105三者求和是11.212。概率分别是0.659、0.242、0.099。如果真实类别是第0类那么交叉熵损失就是-ln0.659≈0.417。这里底数用e也就是自然对数因为PyTorch里的CrossEntropyLoss默认使用ln。如果你用log2计算结果会变成约0.602 bit虽然数值变了但梯度方向完全一致不会影响训练效果。实际工程里我们统一用自然对数即可。3.3 为什么别手写softmaxlog再做交叉熵理论上你可以先算softmax再取log再算损失。但工程上强烈不建议这么做原因有两个。第一个是数值稳定性。如果某个类别的得分特别低softmax算出的概率可能是一个极其小的数比如1e-12取log后得到约-27.6这还勉强能算。但如果概率下溢变成了0.0log(0)就是-inf梯度直接变成NaN训练直接崩掉。第二个在反向传播时的好处。把softmax和交叉熵放在一起梯度会得到一个极其简洁的形式。对单分类场景softmax输出的概率q_i和真实one-hot标签y_i交叉熵损失对logits的偏导数居然是∂L/∂z_i q_i - y_i这个结论漂亮到让人起鸡皮疙瘩。当某个类别是正确类别时y_i1梯度是q_i-1当某个类别不是正确类别时y_i0梯度就是q_i。所以模型会推动正确类别概率升高、错误类别概率降低而这种推动力的大小直接取决于误差误差越大梯度越明显收敛自然就快。如果你把softmax和交叉熵拆开手写中间经过概率值时梯度计算就会引入额外的除法项数值上更容易放大误差。因此PyTorch直接提供了融合好的torch.nn.CrossEntropyLoss它输入的是原始logits你在使用时就不要再手动加softmax了。4. 完整计算示例从手算到代码验证4.1 信息熵的手算三分类天气我们先用一个生活化例子练练手。假设某地天气概率分布为晴0.6、雨0.3、雪0.1计算这个分布的信息熵。H(P) -0.6*ln0.6 - 0.3*ln0.3 - 0.1*ln0.1 -0.6*(-0.5108) - 0.3*(-1.2040) - 0.1*(-2.3026) 0.3065 0.3612 0.2303 0.8980 nat如果用log2计算就是0.8980/0.6931≈1.296 bit。这说明描述这个天气系统平均至少需要1.296位二进制信息。如果天气变成完全均匀的三分类也就是每个概率都是0.3333熵会变成ln3≈1.0986 nat数值变大不确定度更高。这个对比能帮你建立数量感。4.2 交叉熵的手算编码别人的天气继续用天气场景。真实概率分布P[0.6,0.3,0.1]但你的朋友不知道他猜测是Q1[0.4,0.4,0.2]。计算交叉熵H(P,Q1) -0.6*ln0.4 - 0.3*ln0.4 - 0.1*ln0.2 -0.6*(-0.9163) - 0.3*(-0.9163) - 0.1*(-1.6094) 0.5498 0.2749 0.1609 0.9856 nat朋友的猜测与真实分布有偏差所以交叉熵0.9856大于真实熵0.8980。如果他的猜测变成Q2[0.6,0.2,0.2]交叉熵是H(P,Q2) -0.6*ln0.6 - 0.3*ln0.2 - 0.1*ln0.2 0.3065 0.4828 0.1609 0.9502 nat虽然还是大于0.8980但比Q1更接近真实分布。你会发现交叉熵越低代表预测分布Q和真实分布P越接近。这正是分类模型训练的目标让预测分布不断朝真实标签分布靠近直到交叉熵小到可以接受。4.3 分类模型三分类损失logits实战现在看一个真实的模型输出。假设三分类任务某样本真实标签是第2类模型最后一层输出的logits是[1.2, -0.5, 2.8]。先算softmaxexp(1.2)3.3201exp(-0.5)0.6065exp(2.8)16.4446分母3.32010.606516.444620.3712三个类别的预测概率q03.3201/20.37120.1630q10.6065/20.37120.0298q216.4446/20.37120.8072正确类别是索引2所以损失L -ln(0.8072) -(-0.2141) 0.2141这个0.2141就是该样本的交叉熵损失。如果用log_softmax来算会更直接log(q2) z2 - ln(20.3712) 2.8 - 3.0144 -0.2144取负号得到0.2144。因为四舍五入略有误差正常现象。从损失值0.2141可以看出模型对这个样本已经比较有把握了。如果logits变成[1.2, -0.5, 0.3]重新算一下exp(0.3)1.3499分母5.2765q21.3499/5.27650.2558损失变成-ln0.2558≈1.3639损失瞬间大了六倍多。这就是交叉熵对错误预测的“重拳出击”。4.4 PyTorch代码验证几行代码复现结果光手算不过瘾我们直接上代码验证。用PyTorch验证上面的三分类例子import torch import torch.nn.functional as F logits torch.tensor([1.2, -0.5, 2.8], requires_gradTrue) target torch.tensor([2]) # 方式一官方CrossEntropyLoss loss F.cross_entropy(logits.unsqueeze(0), target) print(loss.item()) # 输出约0.2142 # 方式二手动复现 probs F.softmax(logits, dim0) manual_loss -torch.log(probs[2]) print(manual_loss.item()) # 输出约0.2142 # 方式三验证梯度 loss.backward() print(logits.grad) # 梯度约[0.1630, 0.0298, 0.8072-1-0.1928]第三条输出对应之前说的∂L/∂z_i q_i - y_i。真实标签在索引2所以第2类的梯度是0.8072-1-0.1928其他两类梯度就是它们的概率。验证一下正确类别的概率不是1所以被往下压错误类别的概率大于0所以被往上抬。一次反向传播就把每个logits的调整方向说得明明白白。如果你是在训练批量数据target要传入一个一维张量里面是每个样本的类别索引不要传one-hot向量。CrossEntropyLoss内部会自动处理成one-hot形式传错了反而麻烦。5. 实操中绕不开的几个坑5.1 标签平滑one-hot太绝对容易让模型过度自信上面我们一直在用one-hot标签它表达的信息是“正确类别概率为1其他都是0”。但真实世界的数据往往有噪声有些样本可能标错了有些样本本身就模棱两可。如果用one-hot硬逼模型把正确类别的概率逼近1很容易导致过拟合模型会越来越自信对噪声越来越敏感。解决办法是标签平滑。基本思路是把硬标签往里掺一点“均匀概率”y_smooth (1 - ε) * y_onehot ε / K其中K是类别数ε是平滑系数常用0.1。假设三分类真实类别是第0类ε0.1原来的标签[1,0,0]就变成[0.9 0.1/3, 0.1/3, 0.1/3] [0.9333, 0.0333, 0.0333]损失不再是单纯逼第0类概率接近1同时还允许模型保留一定的“怀疑空间”因为第1类和第2类也有一点点真实的概率权重。实际操作中标签平滑往往能提升模型泛化能力尤其在大规模分类和蒸馏任务里特别常见。如果模型训练时出现损失持续下降但验证集效果变差可以试试把ε调大一点比如0.1到0.2。5.2 类别不均衡直接用交叉熵会偏袒多数类分类任务里经常遇到某种类别的样本特别多另一种特别少。直接使用交叉熵损失模型会发现只要把所有样本都预测成多数类别总损失就能压得很低。这是交叉熵天然的偏向不是模型笨而是你的目标函数没有告诉它“少数类更重要”。解决思路有两种。第一种是给每个类别加权重让少数类的损失在总损失里占比更大L -w_c * log(q_c)w_c可以设为样本数量反比比如类别总样本数越少w_c越大。PyTorch里直接在CrossEntropyLoss中传入weight参数即可loss_fn nn.CrossEntropyLoss(weighttorch.tensor([0.3, 1.0, 3.0]))第二种是使用focal loss它会在交叉熵前面乘一个调制因子(1-q_c)^γ。当模型对某个样本已经预测得很准q_c接近1时调制因子接近0损失被压缩当模型预测得不准q_c很小时调制因子接近1损失保持较大。这样训练时模型会把注意力集中在那些困难样本上而不是已经被拿捏得差不多的简单样本。这个技巧在目标检测和长尾分类里特别流行。5.3 数值稳定能跑通和永远不崩是两码事前面提过不要手写log(softmax(x))这里再展开说。softmax里的指数运算很容易溢出比如logits里有200exp(200)会直接变成inf整个loss变成nan。工程上会做一个平移操作logsumexp(z) max_z log(Σ exp(z - max_z))把每个logits都减去最大值再做指数运算这样最大项变成exp(0)1其他项都小于1不会溢出。PyTorch的log_softmax底层就是这么实现的CrossEntropyLoss在计算时也会自动走这条路线所以你传logits进去是最安全的。如果你因为好奇自己写了一个-torch.log(F.softmax(logits, dim1))在小数值上可能没毛病一旦遇到极端样本就可能出NaN。我的习惯是能调库就别手搓尤其是这种已经被无数人踩平的基础操作。5.4 损失异常排查先看标签再看学习率训练中如果发现交叉熵损失变成负数第一反应别怀疑公式错了。交叉熵损失理论上是非负的在单标签场景下它等于-log(q_c)其中q_c在(0,1)之间所以结果是正数。出现负值多半是target传出的标签索引超出了类别数或者误传了one-hot向量导致索引错乱。先检查标签张量的最大值是否小于类别数。如果损失变成NaN原因不外乎几个学习率太大导致梯度爆炸logits里出现了inf标签中有负数。大模型场景下偶尔也会看到损失正常但评估指标原地踏步这时我会怀疑是标签平滑没有作用到验证阶段或者数据增强导致标签噪声分布与平滑假设不符。排查顺序一般先打印一次前向loss再检查梯度的范数最后看logits的统计量基本能定位。5.5 最后的个人心得我自己最早学这三个概念时死记过很多遍公式但每次看完就忘。后来发现最有效的记忆方式是给自己讲故事熵是“一个系统平均需要多少信息来描述”交叉熵是“用错误的概率模型去描述真实分布要多花多少信息”交叉熵损失就是“分类模型对正确类别的负对数概率”。故事顺了公式自然就记住了。还有一个小技巧想分享如果你对推导感兴趣可以自己用纸笔算一次梯度。从L-log(q_c)和q_c exp(z_c)/Σexp(z_j)出发用链式法则推一遍看到最后得到q_i - y_i这个结果时那种“原来如此”的爽感比看十篇博客都有用。这也是我写这篇答疑时最想传递的东西别怕公式拆开揉碎就是加减乘除。交叉熵损失真正迷人的地方是它同时串起了信息论、统计学习和数值优化三块内容。你越是深入做模型训练越会发现这个小小的函数几乎无处不在分类、分割、语言模型、知识蒸馏甚至生成模型里都能看到它的影子。把今天这几个概念吃透以后跟别人讨论模型细节时你就不只是“会用”而是真的“懂它为什么这么设计”了。