做分类训练这么多年,Cross Entropy Loss 可以说是我打交道最频繁的损失函数。图像分类、文本多分类、目标检测里的类别分支,模型架构换了一茬又一茬,但最终收敛用的基本都是交叉熵这一套。这篇文章是“损失函数大汇总”系列的第四篇,专门把 Cross Entropy Loss 拆开讲透:公式怎么来的、梯度怎么算、从 Numpy 到 PyTorch 的代码怎么写,以及实际项目里那些容易踩的坑。
如果你只是会用框架里现成的nn.CrossEntropyLoss(),那这篇笔记可以帮你补上最后一块拼图;如果你正要自己实现损失函数,或者在做 YOLO、GAN 这类模型时被 loss 搞糊涂了,那这篇内容应该能让你少走不少弯路。
1. 从零理解交叉熵:分类任务为什么离不开它
1.1 信息量、熵与交叉熵:先搞懂这三个概念
交叉熵这个名字听起来有点唬人,但其实它的根基就三个概念:信息量、熵、KL散度。
先说信息量。一件事越 unlikely,它带来的信息量就越大。天气预报说“明天晴天”,这是废话,信息量很小;如果它说“明天有十级台风”,那这个信息量就大了。数学家把这个直觉定义成:
$$I(x)=-\log p(x)$$
这里 $p(x)$ 是事件发生的概率。概率越大,$I(x)$ 越小;概率越小,$I(x)$ 越大。负号保证了结果非负,而且直接套用了“稀有事件携带更多信息”的直觉。
再说熵。熵就是“所有可能事件信息量的期望”,衡量一个系统内部的不确定性:
$$H(p)=-\sum_{x} p(x)\log p(x)$$
一个极端例子:一个硬币永远正面朝上,那它的熵是 0,因为结果已经完全确定了;一个均匀硬币的熵最大,因为你永远猜不到下一次是正面还是反面。
交叉熵则是把“真实分布 $p$”和“模型预测分布 $q$”放在一起,计算用 $q$ 来编码 $p$ 需要花费的额外信息量:
$$H(p,q)=-\sum_{x} p(x)\log q(x)$$
如果 $q$ 和 $p$ 完全一样,交叉熵就等于熵;如果 $q$ 偏离 $p$ 很多,交叉熵会变得很大。这就是它名字里“交叉”二字的含义:两个分布相互纠缠之后产生的熵。
1.2 分类任务里的交叉熵到底长什么样
碰上一个具体的分类任务,事情会变得非常清晰。假设一张图片里要么是猫、要么是狗、要么是鸟,真实标签是“猫”,用 one-hot 编码就是:
$$y=[1,0,0]$$
模型输出的是经过 softmax 之后的概率预测 $\hat y=[0.7,0.2,0.1]$。那么交叉熵就是:
$$L=-\sum_{c=1}^{3} y_c\log \hat y_c=-1\cdot \log 0.7 - 0\cdot \log 0.2 - 0\cdot \log 0.1 \approx 0.357$$
注意,因为真实标签是 one-hot,其它维度全是 0,乘以 $\log$ 之后全部归零,所以交叉熵最终简化成了:
$$L=-\log \hat y_{\text{true}}$$
也就是“真实类别对应的预测概率,取负对数”。预测越高,损失越小;预测越低,损失越大。而且这个惩罚不是线性的:预测 0.7 时损失是 0.357,预测 0.1 时损失飙到 2.30,这个差距是爆炸式的。这正是我们想要的——模型越是自信地给出错误答案,代价就越惨痛。
二分类是它的特例。此时模型输出一个 sigmoid 值 $\hat y\in(0,1)$,真实标签 $y\in{0,1}$,损失写成:
$$L=-(y\log \hat y+(1-y)\log(1-\hat y))$$
这就是 PyTorch 里BCELoss的本质,也是目标检测那个类别分支里遍地都是的东西。你会在 YOLO 的源码里反复看到BCEWithLogitsLoss,那就是把 sigmoid 和 BCELoss 合在了一起,专门处理数值稳定性。
2. 公式推导:一文理清 Cross Entropy 的来龙去脉
2.1 先从信息熵出发:KL散度到底衡量了什么
很多讲解直接甩出交叉熵公式,然后说“这就是损失函数”,这让人很懵。我还是习惯从 KL 散度说起。
KL散度衡量的是“用分布 $q$ 去近似分布 $p$ 时,信息损失了多少”。定义如下:
$$D_{KL}(p|q)=\sum_{x} p(x)\log\frac{p(x)}{q(x)}=\underbrace{-\sum_x p(x)\log q(x)}{H(p,q)}-\underbrace{\left(-\sum_x p(x)\log p(x)\right)}{H(p)}$$
所以 KL 散度 = 交叉熵 − 熵。它能写成一个绝对清晰的公式,而且永远非负,只有当 $p=q$ 时取 0。
我拿一个生活化的例子算给你看。假设真实分布 $p$ 是“今天出门 30% 穿雨衣、70% 带伞”,你的模型却认为 $q$ 是“50% 穿雨衣、50% 带伞”。KL 散度就是:
$$D_{KL}(p|q)=0.3\log\frac{0.3}{0.5}+0.7\log\frac{0.7}{0.5}\approx 0.3\times(-0.51)+0.7\times0.337\approx 0.083$$
这个 0.083 本身没有绝对意义,只有相对意义——模型逼近得越差,这个值越大。这才是关键:在分类任务里,我们根本不关心绝对数字,只关心怎么把它降到最低。
2.2 从KL散度到交叉熵:损失函数是怎么一步步落地的
放到分类任务里,真实分布 $p$ 就是 one-hot 的标签向量,比如 $p=[0,1,0]$。这种情况下 $H(p)=-(0\log0+1\log1+0\log0)=0$,于是:
$$D_{KL}(p|q)=H(p,q)-H(p)=H(p,q)$$
也就是说,在分类任务中最小化 KL 散度,等价于最小化交叉熵。这就是为什么损失函数里只写了交叉熵,而不写 KL 散度——one-hot 标签已经把熵这一项归零了。
另一种等价的推导路径是从最大似然估计出发。对一组样本 $(x_i,y_i)$,我们想让模型输出的概率分布最大可能地产生这些真实标签:
$$\hat\theta=\arg\max_\theta \prod_i \hat y_{i,;y_i}$$
取负对数,最大化乘积变成最小化求和:
$$L=-\frac{1}{N}\sum_{i=1}^{N}\log \hat y_{i,;y_i}=\frac{1}{N}\sum_{i=1}^{N}-\log \hat y_{i,;y_i}$$
这就是负对数似然(NLL)。再把它写成 one-hot 的完整矩阵形式:
$$L=-\frac{1}{N}\sum_{i=1}^{N}\sum_{c=1}^{C} y_{i,c}\log \hat y_{i,c}$$
推到这一步,交叉熵损失函数的公式就算正式建立了。我建议你亲手从信息量、熵、KL散度、最大似然这两条路各推一遍,推完之后再看任何框架的文档都不会觉得那个公式是天上掉下来的。
2.3 softmax与交叉熵的梯度推导:为什么这对组合这么丝滑
公式有了,接下来看梯度。这里可以说是整个交叉熵体系里最优雅的部分,也是面试几乎必考的推导题。
softmax 把 logits $z$ 变成概率:
$$\hat y_j=\frac{e^{z_j}}{\sum_{k}e^{z_k}}$$
先求 softmax 的雅可比矩阵。对于第 $j$ 个输出对第 $k$ 个输入求偏导,要分两种情况:
- 当 $j=k$ 时:$\dfrac{\partial \hat y_j}{\partial z_k}=\hat y_j(1-\hat y_j)$
- 当 $j\neq k$ 时:$\dfrac{\partial \hat y_j}{\partial z_k}=-\hat y_j\hat y_k$
可以合并成紧凑写法:
$$\frac{\partial \hat y_j}{\partial z_k}=\hat y_j(\delta_{jk}-\hat y_k)$$
这里的 $\delta_{jk}$ 是克罗内克 delta:$j=k$ 时为 1,否则为 0。这个式子的来源就是商数法则,你自己验证一遍就知道。
现在看单个样本的交叉熵,下标 $i$ 省略:
$$L=-\sum_{c=1}^{C} y_c\log \hat y_c$$
对第 $k$ 个 logit 求偏导:
$$\frac{\partial L}{\partial z_k}=-\sum_{c=1}^{C} y_c\cdot\frac{1}{\hat y_c}\cdot\frac{\partial \hat y_c}{\partial z_k}$$
代入 softmax 的导数:
$$\frac{\partial L}{\partial z_k}=-\sum_{c} y_c\cdot\frac{1}{\hat y_c}\cdot\hat y_c(\delta_{ck}-\hat y_k)=-\sum_{c} y_c(\delta_{ck}-\hat y_k)$$
展开后分成两项:
$$\frac{\partial L}{\partial z_k}=-y_k+\hat y_k\sum_{c}y_c$$
由于真实标签是 one-hot,$\sum_c y_c=1$,最终得到:
$$\frac{\partial L}{\partial z_k}=\hat y_k-y_k$$
这个结果干净得让人舒服:交叉熵配合 softmax 的反向梯度,就是模型预测概率减去真实标签。如果模型预测 0.7,真实是 1,梯度就是 0.7−1=−0.3,把参数往回拉一点;预测过头了,梯度为正,再把参数压一压。这就是为什么分类网络的反向传播这么稳定。
对比一下,如果分类用 MSE 配 softmax,梯度里会带 $\hat y(1-\hat y)$ 这个饱和项,预测接近 0 或 1 时梯度趋近于 0,训练基本卡死。交叉熵那个 $\frac{1}{\hat y}$ 和 softmax 导数的分子正好互相约掉了,这就是这套组合在数学上“丝滑”的本质原因。
3. 代码实现:从零手写 CrossEntropyLoss
3.1 前向传播:Numpy 实现与数值稳定技巧
理论推导再漂亮,最终还是要落到代码里。我从最朴素的 Numpy 手写版本开始。
先写 softmax,这里就涉及前面说的数值稳定性问题。$e^{z}$ 在 $z$ 很大时直接爆炸,变成 inf;解决办法是让每个 logit 先减去同一行的最大值。这个操作不改变概率结果——因为分子分母都乘了 $e^{-m}$,但能确保指数里的数不会超过 0。
import numpy as np def softmax(logits, axis=-1): # 减去最大值,防止 exp 溢出 m = np.max(logits, axis=axis, keepdims=True) exp_x = np.exp(logits - m) return exp_x / np.sum(exp_x, axis=axis, keepdims=True)前向损失:
def cross_entropy_loss(logits, labels, epsilon=1e-12): """ logits: 网络输出(未经 softmax),shape (N, C) labels: one-hot 标签,shape (N, C) """ probs = softmax(logits) # 裁剪到 (epsilon, 1-epsilon),防止 log(0) probs = np.clip(probs, epsilon, 1.0 - epsilon) N = logits.shape[0] loss = -np.sum(labels * np.log(probs)) / N return loss很多人会疑惑,为什么 logits 要传给 cross_entropy 而不是直接传 softmax 后的结果。最基本的理由是数值稳定性:直接把 softmax 之后的值放进np.log,概率一旦接近 0,log 直接给出负无穷。而如果对 logits 做log_softmax,公式就是:
$$\log\hat y_j=z_j-\log\sum_{k}e^{z_k}$$
这里的 $\log\sum e^{z_k}$ 有个专门的算法叫 logsumexp,它能用一种非常稳的方式计算,避免中间过程出现 inf。包括 PyTorch 在内的框架,官方实现全都是走这条路线。
3.2 反向传播:手写梯度并验证正确性
反向传播就一行代码的事,正是因为上面推导出来的 $\hat y-y$:
def cross_entropy_grad(logits, labels): probs = softmax(logits) # 梯度 = (预测概率 - 真实标签) / batch_size return (probs - labels) / logits.shape[0]这里除以 batch_size 是因为前向损失里求了平均,梯度也要跟着平均。如果前向用的是 sum,那反向就不需要做这步除法。
我建议你写完这两个函数之后,用 PyTorch 的自动求导验一次梯度:
import torch import torch.nn.functional as F # 手动算梯度 logits_np = np.array([[1.0, 2.0, 0.5], [0.1, 0.2, 3.0]]) labels_np = np.array([[1.0, 0.0, 0.0], [0.0, 0.0, 1.0]]) manual_grad = cross_entropy_grad(logits_np, labels_np) # 用 PyTorch 验证 logits_t = torch.tensor(logits_np, requires_grad=True) labels_t = torch.tensor(labels_np) loss_t = F.cross_entropy(logits_t, labels_t.argmax(dim=1)) loss_t.backward() print("手动梯度:", manual_grad) print("PyTorch梯度:", logits_t.grad.numpy())两次结果应该完全一致。我第一次自己实现的时候,手写 softmax 导数时把下标搞反了一个方向,结果梯度和框架结果差了三位数,查了半天。强烈建议你做任何自定义损失函数时,都第一时间用自动微分跑一次对拍验证。
3.3 PyTorch 实操:三种写法对比与官方 API 的正确姿势
在实际项目里我们不会真的手写 Numpy 版本,主要是理解原理。PyTorch 里下面三种写法效果完全等价,但踩坑程度完全不一样:
import torch import torch.nn as nn import torch.nn.functional as F logits = torch.randn(8, 10) # 8 个样本,10 个类别 target = torch.randint(0, 10, (8,)) # class index 形式 # 方式一:直接用官方 CrossEntropyLoss,内部会做 log_softmax criterion = nn.CrossEntropyLoss() loss1 = criterion(logits, target) # 方式二:log_softmax + NLLLoss loss2 = F.nll_loss(F.log_softmax(logits, dim=-1), target) # 方式三:手动写出 numpy 版本逻辑 probs = F.softmax(logits, dim=-1) loss3 = F.nll_loss(torch.log(probs + 1e-12), target) print(loss1.item(), loss2.item(), loss3.item())重点提醒:nn.CrossEntropyLoss的 target 接收的是class index(形状(N,)、值范围 0~C−1 的整数张量),不是 one-hot 向量。如果你手头只有 one-hot,得先target.argmax(dim=1)转成索引,或者用target.to(torch.float32)配合内部的标签平滑选项。
还有一个经常被忽略的坑:nn.CrossEntropyLoss()的输入 logits 不需要手动做 softmax,你把已经 softmax 后的概率传进去,反而是在做二次 softmax,损失曲线会变得非常奇怪,而且很难排查。
4. 盘点损失函数的典型应用场景
4.1 目标检测中的交叉熵:YOLO 系列损失函数的变体演进
我最早被交叉熵搞晕,就是在折腾 YOLO 损失函数的时候。看起来 YOLO 里 div 的 loss 很复杂,但拆开看全是交叉熵的影子。
YOLOv1 还比较原始,类别预测用的是简单的平方误差。到 YOLOv3 开始,类别分支就换成了BCEWithLogitsLoss——因为多标签分类里一张图可能同时有“行人”和“自行车”两个类别,每个类别独立做二分类,所以每个输出节点配一个 sigmoid 加 BCE,而不是所有类别共享一个 softmax。这个改动非常重要,它让模型不再被迫在所有类别里选一个。
YOLOv5 和 v8 的损失函数里,分类分支继续用 BCE,另加置信度分支(objectness)也是 BCE。如果你画出 YOLOv8 的 loss 曲线,会看到cls_loss那条线在训练初期快速下降,中后期变得平坦——这就是交叉熵已经收得差不多了的信号。
在目标检测里,正负样本极度不平衡的问题很突出,于是交叉熵被加上了一个调制因子,这就是 Focal Loss:
$$FL(p_t)=-(1-p_t)^\gamma\log(p_t)$$
当样本被模型正确分类且概率很高时,$(1-p_t)^\gamma$ 接近 0,损失被压低;而困难样本 $p_t$ 小,调制因子接近 1,损失几乎不变。这个操作本质就是给交叉熵做一个“困难样本加权”,是类别不平衡场景下交叉熵最成功的改进之一。
4.2 多分类评估闭环:从混淆矩阵到逐类指标
训练的时候用交叉熵一路优化,那模型到底表现如何?单看总损失是远远不够的,还要结合多分类混淆矩阵做逐类分析。
import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix, classification_report import torch import torch.nn.functional as F # 假设 logits 形状为 (N, C),target 为 (N,) preds = logits.argmax(dim=1).cpu().numpy() cm = confusion_matrix(target.cpu().numpy(), preds) sns.heatmap(cm, annot=True, fmt="d", cmap="Blues") plt.xlabel("Predicted") plt.ylabel("True") plt.show() print(classification_report(target.cpu().numpy(), preds, digits=4))这里有个经验:如果总准确率很高,但混淆矩阵里某个对角线元素特别低,说明模型在某个困难类别上完全没学好,交叉熵虽然一直在降,但那只是把容易类别的置信度拉高了。这种情况不要盲目加正则,先回看这个类别的样本数量是不是太少,再决定是加权交叉熵还是做数据增强。
4.3 生成对抗网络里的交叉熵:BCE的另一种打开方式
生成对抗网络的损失函数也是交叉熵的变形。原始 GAN 里,判别器是一个二分类器,判断输入是真实图片还是生成图片,损失就是标准的 BCE:
$$\min_G\max_D\ \mathbb{E}{x}[ \log D(x) ] + \mathbb{E}{z}[ \log(1-D(G(z))) ]$$
判别器 $D$ 尽可能给真实图高分、给生成图低分;生成器 $G$ 则想让 $D$ 对自己的输出给出高分。这个对抗过程里,交叉熵的“决策边界”特性被发挥得淋漓尽致——它天然鼓励置信度往 0 或 1 两极分化,因此 GAN 早期容易出现判别器饱和、梯度消失的问题。后来各种改进如 WGAN、LSGAN,本质上都在试图换掉或者改造交叉熵的二分类实现。
理解这一点对调试 GAN 很有帮助:如果你发现判别器 loss 迅速掉到接近 0,生成器 loss 怎么都涨不动,那大概率是交叉熵饱和了,这时候可以考虑给标签加噪声(one-sided label smoothing),让判别器不要那么极端。
4.4 标签平滑与温度缩放:工程上让交叉熵更好用
直接用交叉熵训练,模型很容易出现过度自信。一个 1000 类的分类任务里,如果某个样本真实标签是 5,模型非要给这个类别打出 0.999 的概率,训练其实还没学会泛化,只是在死记硬背。标签平滑(Label Smoothing)就是专门治这个的:
def smooth_one_hot(labels, num_classes, smoothing=0.1): # 真实类别概率 = 1 - smoothing,其余类别均分 smoothing conf = 1.0 - smoothing smooth_labels = torch.full((labels.size(0), num_classes), smoothing / (num_classes - 1)) smooth_labels.scatter_(1, labels.unsqueeze(1), conf) return smooth_labels把原本的[0, 0, 1, 0]变成[0.033, 0.033, 0.9, 0.033],模型就不需要再追求那个“1 的极致”,训练后期更稳定。这个技巧我在大规模分类任务里几乎必开。
另一个工具是温度缩放。把 softmax 里的 logits 除以一个温度 $T$:
def tempered_softmax(logits, T=1.0): return F.softmax(logits / T, dim=-1)$T>1$ 时概率分布变得更平滑,所有类别的概率都往中间靠,这原本是蒸馏(Knowledge Distillation)里的核心操作,让教师模型输出软标签给学生模型学。顺便说一句,如果你把软标签拿去做交叉熵,实际上就是在最小化两个分布的 KL 散度,这个角度比看公式更直观。
5. 常见问题与避坑实录
5.1 logits 与 softmax:数值稳定性里最容易翻车的点
我见过太多新手在损失函数这块摔跤,最典型的就是“先手动 softmax,再手动 log,再做 cross entropy”。这条路几乎必然碰到 inf 或者 nan,原因就是log(0)和exp(大数)同时出现。
框架内部的log_softmax用的是合并策略,指数、对数、减法都在同一个 stabilize 的逻辑下完成。你自己写的时候务必要遵循两条铁律:
- 永远不要对 softmax 之后的概率直接取 log,除非你做了 clip;
- 永远不要分开算
log(softmax(x)),而是要算x - logsumexp(x)。
下面这个 numpy 函数演示了 logsumexp 的正确打开方式:
def log_softmax_stable(logits): m = np.max(logits, axis=-1, keepdims=True) return logits - m - np.log(np.sum(np.exp(logits - m), axis=-1, keepdims=True))5.2 类别不平衡与噪声标签:加权交叉熵之外的几条路
类别不平衡是分类任务里的常客,最简单的方法是给nn.CrossEntropyLoss传类别权重:
weights = torch.tensor([0.1, 1.0, 5.0]) # 越少样本的类别权重越大 criterion = nn.CrossEntropyLoss(weight=weights)但权重怎么设是个经验活,设大了反而让模型在少数类上过拟合,产生一堆假阳。我用过之后觉得,不如先用混淆矩阵看看哪些类别在互相混淆,再针对性地做数据增强或采样更有效。
噪声标签对交叉熵则是另一回事。交叉熵在底层逻辑上鼓励模型去拟合所有标签,包括错误标签,所以训练后期 loss 会一直缓慢上升。一个相对简单的变通做法是做标签截断:预测概率超过某个阈值(比如 0.9)的样本,直接把标签改成预测值,切断噪声影响。这个操作在真实数据上比调正则参数还管用。
下面整理一个速查表,都是我实际踩过的问题:
| 现象 | 可能原因 | 排查方向 |
|---|---|---|
| loss 初始就是 nan | logits 里有 inf,或标签索引越界 | 检查输入是否有极大值,target 是否在 0~C-1 范围 |
| 多分类 loss 异常低 | 传了 softmax 后的值,导致二次 softmax | 确认传给 CrossEntropyLoss 的是 logits |
| 训练后期 loss 不降 | 标签有噪声或类别不平衡 | 尝试标签平滑、加权、困难样本挖掘 |
| 权重不平衡时效果反而变差 | 权重设置过大 | 画混淆矩阵,回到数据层面解决 |
| 梯度爆炸 | 学习率太大或 softmax 数值不稳定 | 降低 lr,检查 logits 是否被 normalize |
5.3 训练曲线怎么看:从 loss 收敛判断模型是否在正常学习
训练模型不能只把 loss 打印出来就算完。我养成的习惯是每轮记录 train loss 和 validation loss,训练完直接画图,能够快速判断模型处于什么状态。
基于 YOLO 这类检测任务的习惯,我会把 logger 里的几个 loss 分别画出来。以分类任务为例,画法类似:
import matplotlib.pyplot as plt epochs = list(range(1, len(train_loss) + 1)) plt.plot(epochs, train_loss, label="train CE loss") plt.plot(epochs, val_loss, label="val CE loss") plt.yscale("log") plt.xlabel("epoch") plt.ylabel("loss") plt.legend() plt.grid(True) plt.show()如果 train loss 持续下降而 val loss 在某个 epoch 后反弹,那就是过拟合信号,交叉熵把训练集背下来了;如果两条线同时不降,看看是不是学习率太低或者模型容量不够。如果你发现曲线非常抖,我一般会检查 batch size 是不是太小,或者数据增强是不是太狠了。
把 loss 曲线和混淆矩阵、逐类指标放在一起看,比单独盯一个数字有效得多。
最后说点我自己的体会。Cross Entropy Loss 看起来简单,但真正用熟它需要完成三步:手动过一遍推导、手写一遍 Numpy 版本、再回到框架里跑通一次官方的接口。我最早从框架里的现成损失函数切换到手写版本时,梯度对拍花了一个下午,最后发现是 softmax 导数里下标方向写反了。排完之后我反而觉得特别值——正是那次折腾,才让我在之后看 YOLO 的损失函数、改 GAN 的判别器、调标签平滑时,都能直接回到最原始的公式找答案。如果你也正在被某个损失函数困扰,不妨先放下框架,把前向和反向各手写一遍,很多坑自然会浮出水面。