AI 里最常做的那一次改写
这一章的改写你每天都在用,只是它藏在框架里。softmax、交叉熵、log 概率、注意力分数——这一整条链上的每一环都会溢出或者下溢,而它们全部靠同一个技巧活着:先减掉最大值。搞懂它顺便能解释一个常见困惑:为什么 PyTorch 要单独提供 log_softmax,而不让你写 log(softmax(x))。
一个三分类模型输出了三个 logit:
z = [1000, 1001, 1002]
这三个数本身完全正常——只是模型比较自信,或者你忘了做 layer norm。照 softmax 的定义算:
const e = z.map(Math.exp); const s = e.reduce((a, b) => a + b, 0); const p = e.map(v => v / s);
问:p 是什么?(数学上的正确答案是 [0.0900, 0.2447, 0.6652]——只和 logit 之间的差有关,和它们有多大无关)
顺便预测:logit 要多大才会出问题?1000?100?还是更小?
格子的另一头
前面十章讲的抹零,都发生在刻度太粗的地方。这一章讲的是另外两种失败:数掉到格子外面去了。
exp(709) = 8.218407461554972e+307 ← 还在
exp(710) = Infinity ← 出界了
ln(最大 double) = 709.782712893384
所以:只要有一个 logit 超过 709.78,exp 就溢出。
三个 Infinity 相加还是 Infinity。
Infinity / Infinity = NaN。
答案是 D。
另一头也一样危险,而且更隐蔽:
exp(−745) ≈ 5e−324 ← 最小的次正规数,只剩一位有效数字 exp(−746) = 0 ← 下溢,安静地变成 0 下溢不会给你 NaN,不会报错。它给你一个 0, 然后 log(0) = −Infinity,然后 −Infinity × 0 = NaN, NaN 传染到梯度里,训练崩掉——而崩的地方离出问题的地方隔了三层。
回答第二问:709.78 是溢出的门槛,但真正开始出问题远比这早。logit 到 100 时,exp(100) ≈ 2.7×10⁴³,还没溢出,可如果同一批里还有个 exp(−100) ≈ 3.7×10⁻⁴⁴,两者相加时小的那个会被完全吃掉(第 2 章:相对差 10⁸⁷,远超 double 的 10¹⁶ 分辨力)。信息在溢出之前很久就开始丢了。
改写:减掉最大值
softmax 有一条恒等式,可以一眼看出来:
softmax(z)ᵢ = exp(zᵢ) / Σⱼ exp(zⱼ)
分子分母同时除以 exp(m)(m 是任意常数):
= exp(zᵢ − m) / Σⱼ exp(zⱼ − m)
对任意 m 都成立。那就取 m = max(z)。
取最大值有两个好处,正好各堵一头:
- 最大的那一项变成
exp(0) = 1⇒ 永远不会溢出。 - 分母至少是 1 ⇒ 除法永远不会除以 0。
// 这么写:直接照定义 const e = z.map(Math.exp); const s = e.reduce((a, b) => a + b, 0); const p = e.map(v => v / s); // [NaN, NaN, NaN] // 改成:先减最大值 const m = Math.max(...z); // ★ 唯一的改动 const e = z.map(v => Math.exp(v - m)); const s = e.reduce((a, b) => a + b, 0); const p = e.map(v => v / s); // [0.0900305732, 0.244728471, 0.665240956]
# 减最大值之后的中间量 exp(1000 − 1002) = 0.135335283 exp(1001 − 1002) = 0.367879441 exp(1002 − 1002) = 1.00000000 # 结果 ✓ [0.0900305732, 0.244728471, 0.665240956] ✓ 三项之和 = 0.9999999999999999 ★ 与 numpy 的稳定版 softmax 逐位相同
下溢那一侧会怎样?如果某个 zᵢ − m 是 −1000,exp 给出 0——而这没关系:那一项的真实概率本来就是 e⁻¹⁰⁰⁰ ≈ 10⁻⁴³⁴,它对结果的贡献在 double 里根本表示不出来。下溢成 0 在这里是「正确舍入」,不是错误。
整条链上,凡是「一堆 exp 加起来再取 log」的地方,都要用同一个恒等式:
logsumexp(z) = m + log Σᵢ exp(zᵢ − m), 其中 m = max(z)
这个函数叫 log-sum-exp,是概率计算里最重要的一个数值原语。softmax、交叉熵、log 概率的归一化、隐马尔可夫的前向算法、变分推断里的 ELBO、混合模型的对数似然——全部是它的马甲。
验算:logsumexp([1000,1001,1002]) = 1002.4076059644444。直接照定义算是 log(Infinity) = Infinity。
为什么有 log_softmax
现在看一个更容易被忽略的问题。训练时你需要的其实不是概率,是 log 概率(交叉熵损失 = 负的 log 概率)。最自然的写法是:
loss = -log(softmax(z)[target])
就算 softmax 用的是稳定版,这一行还是有问题。看这个例子:
z = [0, −800, −900] 稳定版 softmax ⇒ [1, ⟨0⟩, ⟨0⟩] ← 后两项下溢成 0(真值是 1e−348 和 1e−391) log(softmax) ⇒ [0, −Infinity, −Infinity] ✗ log_softmax ⇒ [0, −800, −900] ✓ 精确
差别在于:softmax 的输出必须落在 [0,1] 里,这个区间装不下 e⁻⁸⁰⁰;而 log 概率的输出是 −800,double 表示它绰绰有余。先算 softmax 再取 log,等于先把答案压进一个装不下它的格子,再想拿出来。
正确的做法是永远不要中途变回概率:
log_softmax(z)ᵢ = zᵢ − logsumexp(z)
= zᵢ − m − log Σⱼ exp(zⱼ − m)
# 一次 exp 都没有作用在最终结果上,全程待在 log 空间。
这就是为什么每个框架都有 log_softmax、cross_entropy(直接吃 logit,不吃概率)、binary_cross_entropy_with_logits(名字里那个 with_logits 就是这个意思)。它们不是便利函数,它们是数值上唯一正确的写法。
logit 这个词在机器学习里被用得很松,值得正一下:
- 统计学原义:
logit(p) = log(p/(1−p)),是 sigmoid 的反函数。 - 深度学习里的用法:「送进 softmax 之前的那个未归一化分数」。这是个借用,严格说不是同一个东西(softmax 前的分数可以差一个任意常数,logit 不行)。
但这个借用抓住了要害:logit 空间是加性的,概率空间是乘性的。浮点数擅长表示跨越很多数量级的量(因为它存的是指数),却不擅长表示「非常接近 0 或 1 的概率」(因为那要求绝对精度)。所以数值计算里的通则是:能待在 log 空间就别出来。
这一条和《意外》那本书是同一件事的两面——那里说「信息量 = −log p」,这里说「log 空间是概率的正确表示」。信息论选 log 是因为可加,数值计算选 log 是因为不溢出,而这两个理由其实是同一个理由。
三行就能看到 NaN,再三行看到它被修好:
import numpy as np
z = np.array([1000.0, 1001.0, 1002.0])
with np.errstate(over='ignore', invalid='ignore'):
e = np.exp(z); print('朴素 softmax :', e / e.sum()) # [nan nan nan]
m = z.max()
e = np.exp(z - m); print('稳定 softmax :', e / e.sum()) # [0.09 0.2447 0.6652]
print('logsumexp :', m + np.log(np.exp(z - m).sum())) # 1002.4076059644444
# log(softmax) vs log_softmax
z2 = np.array([0.0, -800.0, -900.0])
p = np.exp(z2 - z2.max()); p /= p.sum()
with np.errstate(divide='ignore'):
print('log(softmax) :', np.log(p)) # [0. -inf -inf]
print('log_softmax :', z2 - (z2.max() + np.log(np.exp(z2 - z2.max()).sum())))
# [0. -800. -900.]
print('ln(最大 double) =', np.log(np.finfo(np.float64).max)) # 709.782712893384
python3 lse.py
在线:Google Colab(numpy 预装)。装了 PyTorch 的话,把上面几行和 torch.softmax / torch.log_softmax 对一遍——会逐位相同,因为它们做的是同一件事。
注意力机制。Transformer 里 QKᵀ/√d 的结果直接进 softmax。那个 √d 的除法不只是「让方差归一」的理论考虑,它同时是在压住 logit 的量级——没有它,随着维度增大分数会线性变大,最终撞上 709。而 FlashAttention 之所以能分块计算 softmax 而结果不变,靠的正是「先减最大值」这个恒等式:每块记住自己的最大值和局部和,合并时重新对齐。这一章的技巧是那篇论文的地基。
隐马尔可夫模型与 Viterbi。连乘几百个小于 1 的概率,几十步就下溢到 0。所有实现一律在 log 空间做加法。《对上》那本书里的日语分词器(Viterbi)就是这么写的。
贝叶斯推断。后验 ∝ 似然 × 先验,似然是几万个数据点的概率连乘。不取 log 的话第一百个点就没了。logsumexp 在算证据、算混合模型、算重要性采样权重时到处都是。
推荐系统的召回打分、语言模型的 beam search、对比学习的 InfoNCE——凡是「一堆分数归一化成概率」的地方,都是同一个函数。
「softmax 只和 logit 之间的差有关,所以 logit 本身多大都无所谓。」
前半句在数学上完全正确——这正是「减掉最大值」这个改写成立的理由。但它对计算路径不成立:exp(1000) 这个中间量是实实在在要被算出来的,而它装不下。
这是这本书反复出现的同一个错误形态,第 9 章刚见过一次:「最终结果与 X 无关」不蕴含「计算过程与 X 无关」。方差与均值无关(第 10 章),但一遍法的计算过程与均值极其有关;softmax 与 logit 的绝对大小无关,但朴素实现与它极其有关。
可操作的版本:看到一个「结果不依赖某个量」的性质,去检查这个性质在你的代码里是不是显式成立的。如果它只是「数学上会抵消」,那它在浮点里多半不会抵消——要主动把它消掉,这就是减最大值在做的事。
正确答案是 D:[NaN, NaN, NaN]。溢出门槛是 709.78,但信息在那之前很久就开始丢了。
A 「softmax 只看差值,所以没问题」——数学上对,见上面那一栏。这也是最多人选的答案,因为它用的是正确的数学直觉,只是没跟着代码走一遍。正确的数学直觉在浮点世界里是不够的,你还得走一遍那条路径。 B 「[0.333, 0.333, 0.333]」——这个答案想的是「精度不够,三个数区分不出来了」。这在另一种输入下会真的发生:如果 logit 是[1e17, 1e17+1, 1e17+2],那么第 2 章的问题就来了——ulp(1e17) = 16,那三个 logit 存进 double 之后完全相同,softmax 确实给出 [0.333, 0.333, 0.333]。B 是一个真实的失败模式,只是不在这道题上。
C 「[0, 0, 1]」——这是「softmax 退化成 argmax」的情形,在 [0, −800, −900] 这类输入上确实会发生(下溢)。它和 D 是同一条链上的两头:大的一头给 NaN,小的一头给 0。两头都由同一个改写堵上。
这一章的改写有一条主线和三条推论:
exp(z) / sum(exp(z)) 然后 log(...)
m = max(z); exp(z-m) / sum(exp(z-m)) 永不溢出,永不除零
log_softmax(z) = z - (m + log(sum(exp(z-m)))) 全程待在 log 空间
框架里直接调 cross_entropy(logits, target),不要自己先算 softmax 再算 log
三条推论,都是同一件事的不同说法:
- 概率连乘 → log 相加。连乘会下溢,相加不会。
- 概率相加 → logsumexp。这是 log 空间里唯一麻烦的操作,所以它值得一个专门的函数名。
- 接口设计上:让函数吃 logit,而不是吃概率。把归一化留在函数内部,调用方就没机会写错。这是框架 API 里那些
_with_logits后缀的全部含义。
这一章的一句话
概率的正确表示是 log 概率——不是因为好看,是因为浮点数存的是指数,而概率的信息全在指数里;能不出 log 空间,就别出来。
下一章收掉卷 III,讲一个每天都在写、几乎每次都写错的东西:两个浮点数怎么比大小。a === b 不行,Math.abs(a-b) < 1e-9 也不行——它会把相差 400% 的两个数判成相等。而这一章还会给出一个让人意外的事实:0.1 + 0.2 和 0.3 之间,其实只隔了一格。它们是邻居。