【无标题】

彻底搞懂 PyTorch 的 CrossEntropyLoss:从 logits 到交叉熵

面向有一点 Python、但没有机器学习背景的读者。全文用一个「预测 RNA 碱基」的四分类例子贯穿,所有数字都可复现。

目录


一、先问一个问题:模型为什么不直接给答案

假设我们在做 RNA 序列设计:给定一个三维骨架,模型要在每个位置上判断这里应该是 A、C、G、U 中的哪一个。

最直觉的做法是让模型直接输出一个字母。但这样做训练不起来,原因是:

训练需要知道「错了多少」,而不只是「错没错」。

如果模型只给一个硬邦邦的答案,猜错了你只知道它错了,却不知道该往哪个方向、调多大幅度去改参数。梯度下降需要一个连续的、可微的数值来指示改进方向,而「对/错」是离散的,导数处处为零。

所以我们换一种方式:让模型对每一个候选类别都表个态,给出一组连续的分数。有了连续的分数,才谈得上「离正确答案还差多少」。


二、logits:模型吐出来的原始分数

对四分类问题,模型在每个位置吐出 4 个数:

复制代码
      A      C      G      U
   [ 2.0,   0.1,   0.1,   0.1 ]

这组数就叫 logits。

关于它,需要明确三点:

  1. 它不是概率。 可以是任意实数,可以为负,加起来也不等于 1。
  2. 只有相对大小有意义。 谁大就说明模型更倾向谁。
  3. 绝对数值不重要。 后面会看到,给所有 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 内部

四条最该记住的:

  1. logits 是原始分数,不是概率,模型最后一层不要接 softmax
  2. 标签是整数下标,dtype 是 long,不是 one-hot
  3. 高维输入时类别维必须在 dim=1 ,不确定就摊平成 [N, C] 和 [N]
  4. ln ⁡ C \ln C lnC 是瞎猜的损失,四分类是 1.386,二十分类是 2.996------训练时拿它当标尺

本文所有数值均可用文中代码复现。如有错漏欢迎指正。

相关推荐
迅猛龙办公室42 分钟前
Python实现简单的人名对话
python
阿坨1 小时前
firestart:一行命令启动你的日常应用和网页
python·pypi·cli·click
老歌老听老掉牙1 小时前
麻花钻切屑形态演变的力学机制与临界条件分析
python·算法·钻头
happylifetree2 小时前
Python18(补充):练习
python
bigdata-余建新2 小时前
week2
人工智能·pytorch·深度学习
weixin_416667962 小时前
银河麒麟V10看门狗试验-1
网络·chrome·python
古城小栈3 小时前
Pydantic 从入门到实践全讲解
python
happylifetree3 小时前
Python18:核心语法-数据存储与运算-运算符-算术运算符
python
打工仔折腾 AI3 小时前
从Attention到BERT:双向预训练语言模型到底解决了什么问题
人工智能·后端·python·深度学习·语言模型·bert