彻底搞懂 PyTorch 的 CrossEntropyLoss:从 logits 到交叉熵
面向有一点 Python、但没有机器学习背景的读者。全文用一个「预测 RNA 碱基」的四分类例子贯穿,所有数字都可复现。
目录
- 一、先问一个问题:模型为什么不直接给答案
- 二、logits:模型吐出来的原始分数
- 三、softmax:从分数到概率
- 四、负对数:从概率到损失
- 五、合起来就是交叉熵
- 六、完整代码验证
- 七、形状契约
- 八、五个常见的坑
- 九、和其它损失函数的关系
- 十、参数速查
- 十一、一张表收尾
一、先问一个问题:模型为什么不直接给答案
假设我们在做 RNA 序列设计:给定一个三维骨架,模型要在每个位置上判断这里应该是 A、C、G、U 中的哪一个。
最直觉的做法是让模型直接输出一个字母。但这样做训练不起来,原因是:
训练需要知道「错了多少」,而不只是「错没错」。
如果模型只给一个硬邦邦的答案,猜错了你只知道它错了,却不知道该往哪个方向、调多大幅度去改参数。梯度下降需要一个连续的、可微的数值来指示改进方向,而「对/错」是离散的,导数处处为零。
所以我们换一种方式:让模型对每一个候选类别都表个态,给出一组连续的分数。有了连续的分数,才谈得上「离正确答案还差多少」。
二、logits:模型吐出来的原始分数
对四分类问题,模型在每个位置吐出 4 个数:
A C G U
[ 2.0, 0.1, 0.1, 0.1 ]
这组数就叫 logits。
关于它,需要明确三点:
- 它不是概率。 可以是任意实数,可以为负,加起来也不等于 1。
- 只有相对大小有意义。 谁大就说明模型更倾向谁。
- 绝对数值不重要。 后面会看到,给所有 logits 同时加一个常数,最终结果完全不变。
logit 这个词源自统计学里的「对数几率」(log-odds),但在深度学习的语境下不必纠结词源。记住一句话就够了:logits 是还没有被转换成概率的原始分数。
那为什么不干脆让模型直接输出概率呢?因为要求神经网络的最后一层输出「四个非负数、且恰好加起来等于 1」是一个很别扭的约束。更省事的做法是:让它随便吐四个实数,然后用一个固定的、可微的公式把它们换算成概率。
那个公式就是 softmax。
三、softmax:从分数到概率
softmax 做两件事:先取指数把所有数变成正数,再除以总和让它们加起来等于 1。
p i = e z i ∑ j e z j p_i = \frac{e^{z_i}}{\sum_{j} e^{z_j}} pi=∑jezjezi
用上面那组 logits 实际算一遍:
| 类别 | logit z z z | e z e^z ez | 除以总和 | 概率 |
|---|---|---|---|---|
| A | 2.0 | 7.389 | 7.389 / 10.705 | 0.6903 |
| C | 0.1 | 1.105 | 1.105 / 10.705 | 0.1032 |
| G | 0.1 | 1.105 | 1.105 / 10.705 | 0.1032 |
| U | 0.1 | 1.105 | 1.105 / 10.705 | 0.1032 |
| 总和 10.705 | 合计 1.0000 |
有两个现象值得注意。
第一,指数会放大差距。 logit 上 A 只比 C 高 1.9,概率上却变成了约 6.7 倍。softmax 里是指数关系,所以 logits 上不大的差距,在概率上会被拉得很开。
第二,softmax 具有平移不变性。 给所有 logits 同时加 100,结果一模一样:
python
softmax([2.0, 0.1, 0.1, 0.1]) == softmax([102.0, 100.1, 100.1, 100.1]) # True
这正是前面说「绝对数值不重要」的原因。工程上也靠这条性质做数值稳定:实现时先减去最大值再取指数,避免 e z e^{z} ez 溢出。
给生物背景读者的一个类比:softmax 的输出其实就是一列 PWM(位置权重矩阵)。模型在说「这个位置我认为 69% 是 A,各 10% 是 C/G/U」。如果你熟悉序列 logo,这一步的输出就是画 logo 用的那种概率分布。
四、负对数:从概率到损失
现在模型给出了一个概率分布,我们需要把它变成一个损失------一个「越小越好」的数。
损失函数要回答的问题只有一个:
真实答案那个类别,模型给了多少概率?
- 真实是 A,模型给了 69% → 还不错 → 损失小
- 真实是 C,模型只给了 10% → 很差 → 损失大
把概率转换成损失,用的是取负对数 : loss = − ln p \text{loss} = -\ln p loss=−lnp。
| 模型给正确答案的概率 | − ln p -\ln p −lnp |
|---|---|
| 100% | 0 |
| 90% | 0.105 |
| 69% | 0.371 |
| 50% | 0.693 |
| 25%(四选一瞎猜) | 1.386 |
| 10% | 2.303 |
| 1% | 4.605 |
| 0.1% | 6.908 |
为什么是 − ln p -\ln p −lnp 而不是 1 − p 1-p 1−p
看表格最后几行。概率 1% 和 10%,在 1 − p 1-p 1−p 看来是 0.99 和 0.90,几乎没有区别;但在 − ln p -\ln p −lnp 看来是 4.61 和 2.30,相差一倍。
「把正确答案判断成几乎不可能」必须付出极大的代价 ------这正是训练所需要的惩罚力度。对数函数在 p → 0 p \to 0 p→0 时趋于无穷,恰好提供了这种性质。
另外, − ln p -\ln p −lnp 不是拍脑袋想出来的。它就是负对数似然(negative log-likelihood):最小化它,等价于最大化模型给训练数据的似然。生物信息里做序列比对打分、PWM 扫描 motif 用的 log-odds 分数,背后是同一套逻辑。
一个必须记住的刻度
C C C 类均匀瞎猜的损失是 ln C \ln C lnC。
| 类别数 C C C | 瞎猜时的损失 |
|---|---|
| 2(二分类) | 0.693 |
| 4(ACGU) | 1.386 |
| 20(氨基酸) | 2.996 |
训练时盯住这个数。看见 loss 从 1.386 往下掉,说明模型开始比瞎猜强;一直卡在 1.386 不动,说明它什么也没学到。这是排查训练问题的第一个判据。
五、合起来就是交叉熵
把前两步合在一起,就是 PyTorch 的 nn.CrossEntropyLoss:
loss = − ln e z y ∑ j e z j = − z y + ln ∑ j e z j \text{loss} = -\ln \frac{e^{z_y}}{\sum_j e^{z_j}} = -z_y + \ln \sum_j e^{z_j} loss=−ln∑jezjezy=−zy+lnj∑ezj
其中 z z z 是 logits, y y y 是真实类别的下标。
「交叉熵」这个名字听起来很唬人。它的一般定义是两个分布 p p p(真实)和 q q q(预测)之间的量:
H ( p , q ) = − ∑ i p i ln q i H(p, q) = -\sum_i p_i \ln q_i H(p,q)=−i∑pilnqi
但在单标签分类里,真实分布 p p p 是一个 one-hot 向量------正确类是 1,其余全是 0。代进去之后,求和里只剩下一项:
H ( p , q ) = − ln q y H(p, q) = -\ln q_y H(p,q)=−lnqy
所以在分类问题中,交叉熵就等于「正确那一类的概率的负对数」,没有别的。 上面那些铺垫,就是它的全部内容。
六、完整代码验证
python
import torch
import torch.nn as nn
# 两个样本,四个类别(A C G U)
logits = torch.tensor([[2.0, 0.1, 0.1, 0.1], # 这一行模型倾向第 0 类
[0.1, 0.1, 3.0, 0.1]]) # 这一行模型倾向第 2 类
y = torch.tensor([0, 2]) # 真实标签,整数下标
loss_fn = nn.CrossEntropyLoss()
print("猜对时 loss =", loss_fn(logits, y).item())
print("猜错时 loss =", loss_fn(logits, torch.tensor([1, 1])).item())
输出:
猜对时 loss = 0.26173...
猜错时 loss = 2.66173...
把黑盒拆开,手算一遍
python
p = torch.softmax(logits, dim=1)
print("概率分布:\n", p)
# tensor([[0.6903, 0.1032, 0.1032, 0.1032],
# [0.0472, 0.0472, 0.8583, 0.0472]])
print("第0个样本,正确类(0)的概率 =", p[0, 0].item()) # 0.6903
print("第1个样本,正确类(2)的概率 =", p[1, 2].item()) # 0.8583
手算 = (-torch.log(p[0, 0]) - torch.log(p[1, 2])) / 2
print("手算 loss =", 手算.item()) # 0.2617
print("库算 loss =", nn.CrossEntropyLoss()(logits, y).item())
逐项对照:
| 样本 | 真实类 | 正确类概率 | − ln p -\ln p −lnp |
|---|---|---|---|
| 0 | 0 | 0.6903 | 0.3711 |
| 1 | 2 | 0.8583 | 0.1528 |
| 平均 0.2617 |
换成错误标签 [1, 1]:
| 样本 | 被当成的类 | 该类概率 | − ln p -\ln p −lnp |
|---|---|---|---|
| 0 | 1 | 0.1032 | 2.2707 |
| 1 | 1 | 0.0472 | 3.0528 |
| 平均 2.6617 |
七、形状契约
这是实际写代码时最容易出错的地方。
基本情形
| 参数 | 形状 | 类型 |
|---|---|---|
input(logits) |
[N, C] |
float |
target(标签) |
[N] |
long(整数下标,不是 one-hot) |
| 返回值 | 标量 | float |
N是样本数,C是类别数- 标签取值范围是
0 ~ C-1 - 标签必须是
torch.long,传 float 会直接报错 - 不要传 one-hot。 PyTorch 要的是下标
高维情形:一个很坑的约定
当输入多于两维时,PyTorch 规定 类别维必须在第 1 维(dim=1):
input : [N, C, d1, d2, ...]
target: [N, d1, d2, ...]
这个约定和大多数人写序列模型的习惯相反。假设你的模型输出是 [batch, 序列长度, 类别数],直接丢进去是错的:
python
logits = torch.randn(8, 100, 4) # [B, L, C] ← 常见的模型输出布局
target = torch.randint(0, 4, (8, 100))
# 错误:会把 L=100 当成类别数,要么报错,要么算出毫无意义的结果
# loss = nn.CrossEntropyLoss()(logits, target)
# 正确写法一:把类别维换到第 1 维
loss = nn.CrossEntropyLoss()(logits.permute(0, 2, 1), target) # [B, C, L]
# 正确写法二:直接摊平成二维(更不容易出错)
loss = nn.CrossEntropyLoss()(logits.reshape(-1, 4), target.reshape(-1))
个人建议用第二种。 摊平成 [N, C] 和 [N] 之后,不存在任何歧义,读代码的人也一眼能看懂。
八、五个常见的坑
坑 1:自己先做了一次 softmax(最常见)
python
# 错误示范
probs = torch.softmax(logits, dim=1)
loss = nn.CrossEntropyLoss()(probs, y) # softmax 被做了两遍
CrossEntropyLoss 内部已经包含 softmax。再手动做一次,相当于对已经归一化的概率再归一化一次,分布会被严重压平:
做一次 softmax: [0.6903, 0.1032, 0.1032, 0.1032]
做两次 softmax: [0.3748, 0.2084, 0.2084, 0.2084] ← 正确类从 0.69 掉到 0.37
对应的 loss: 0.371 → 0.981
这个错误不会报错,训练照常进行,loss 也在下降,只是梯度被压扁、模型学得异常慢。排查起来非常费劲。
规则:模型最后一层直接输出 logits,不要接 softmax。 只有在推理阶段需要看概率时,才单独调用 torch.softmax。
坑 2:标签用了 one-hot
python
y = torch.tensor([[1., 0., 0., 0.], [0., 0., 1., 0.]]) # 不需要这样
y = torch.tensor([0, 2]) # 这样就对了
(补充:PyTorch 1.10 之后 CrossEntropyLoss 也接受「概率型 target」,即和 input 同形状的浮点张量,用于标签平滑或知识蒸馏。但常规单标签分类一律用整数下标。)
坑 3:标签 dtype 不对
标签必须是 torch.long。从 numpy 转过来时特别容易带成 int32 或 float64:
python
y = torch.as_tensor(np_labels, dtype=torch.long) # 显式指定
坑 4:类别数和标签取值对不上
input 的最后一维是 C,标签取值必须落在 [0, C-1]。出现 C 或更大的值会触发越界错误,在 GPU 上报的错往往是一句无关的 device-side assert,极难定位。
建议:训练前先断言一次。
python
assert int(target.max()) < logits.shape[-1], "标签越界了"
assert int(target.min()) >= 0
坑 5:忘了 loss 默认是取平均
reduction='mean' 是默认值,返回的是所有样本损失的平均。如果你想自己按长度加权,或者要做梯度累积,记得用 reduction='sum' 或 'none'。
另外一个容易忽略的细节:当你传了 weight 参数,'mean' 的分母是权重之和,不是样本数 N。
九、和其它损失函数的关系
CrossEntropyLoss = LogSoftmax + NLLLoss
python
# 下面两种写法完全等价
loss1 = nn.CrossEntropyLoss()(logits, y)
log_p = torch.log_softmax(logits, dim=1)
loss2 = nn.NLLLoss()(log_p, y)
NLLLoss 的名字是 Negative Log Likelihood------它只负责「挑出正确类、取负号」,softmax 和取对数要你自己做。CrossEntropyLoss 把三步打包了,而且在数值上更稳定(内部用了 log-sum-exp 技巧),所以优先用它。
二分类:和 BCEWithLogitsLoss 的关系
二分类有两种写法,结果完全等价:
python
import numpy as np
z0, z1 = 0.4, 1.3
# 写法一:当成 2 类的多分类,模型输出 2 个 logits
ce = -np.log(np.exp(z1) / (np.exp(z0) + np.exp(z1))) # 0.341154
# 写法二:模型只输出 1 个 logit(两者之差),用 BCEWithLogitsLoss
bce = np.log1p(np.exp(-(z1 - z0))) # 0.341154
选哪个 :类别互斥(一个样本只能属于一类)用 CrossEntropyLoss;多标签 (一个样本可以同时属于多个类,比如一段序列同时带好几种功能注释)必须用 BCEWithLogitsLoss,对每个标签独立判断。
一句话区分三兄弟
| 损失函数 | 输入 | 适用场景 |
|---|---|---|
CrossEntropyLoss |
logits | 单标签多分类(类别互斥) |
NLLLoss |
log 概率 | 同上,但你已经自己做过 log_softmax |
BCEWithLogitsLoss |
logits | 二分类 / 多标签(类别不互斥) |
十、参数速查
python
nn.CrossEntropyLoss(
weight=None, # [C] 的张量,给每个类别不同权重,用于类别不平衡
ignore_index=-100, # 标签等于这个值的位置直接跳过,常用于 padding
reduction='mean', # 'mean' | 'sum' | 'none'
label_smoothing=0.0, # 标签平滑,缓解模型过度自信(PyTorch 1.10+)
)
实际用得最多的两个:
weight ------ 类别不平衡。 比如某一类样本特别少,可以给它更大的权重。注意前面提过的细节:配合 reduction='mean' 时,分母会变成权重之和。
ignore_index ------ 变长序列的 padding。 做序列任务时,把 padding 位置的标签设成 -100,损失就会自动忽略它们,不用手写 mask:
python
target[pad_mask] = -100
loss = nn.CrossEntropyLoss()(logits, target) # padding 位置不参与计算
label_smoothing 的具体实现约定建议查官方文档确认,这里不展开。
十一、一张表收尾
把整个流程串起来:
| 步骤 | 输入 | 输出 | 由谁负责 |
|---|---|---|---|
| 1. 模型前向 | 特征 | logits [N, C] |
你的模型 |
| 2. softmax | logits | 概率分布 [N, C] |
CrossEntropyLoss 内部 |
| 3. 取出正确类的概率 | 概率 + 标签 | [N] |
CrossEntropyLoss 内部 |
| 4. 取负对数 | 概率 | 损失 [N] |
CrossEntropyLoss 内部 |
| 5. 求平均 | 损失 | 标量 | CrossEntropyLoss 内部 |
四条最该记住的:
- logits 是原始分数,不是概率,模型最后一层不要接 softmax
- 标签是整数下标,dtype 是 long,不是 one-hot
- 高维输入时类别维必须在 dim=1 ,不确定就摊平成
[N, C]和[N] - ln C \ln C lnC 是瞎猜的损失,四分类是 1.386,二十分类是 2.996------训练时拿它当标尺
本文所有数值均可用文中代码复现。如有错漏欢迎指正。