☰
Cross Entropy Loss深度解析:从公式推导到PyTorch实现
2026/10/10 6:44:26 网站建设 项目流程

做分类训练这么多年,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 初始就是 nanlogits 里有 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 的判别器、调标签平滑时,都能直接回到最原始的公式找答案。如果你也正在被某个损失函数困扰,不妨先放下框架,把前向和反向各手写一遍,很多坑自然会浮出水面。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询