损失与稳定性 入门

分类器的最后一层吐出的是分数,而从分数走到损失的每一步,都可能让算术悄悄不再是算术。softmax、平方误差对交叉熵、 88.72 这条指数上限、把它消掉的「减最大值」、log-sum-exp,以及为什么各家框架的损失函数拒收概率。这一页上的每个数字都由旁边那张图在模拟的 float32 里算出来 —— 所以你可以亲手把它推过上限,看它坏掉。

01

分数不是概率

分类器最后一层每个类别吐出一个数。这些数身上没有任何一点像分布。

分类器最后那个线性层,每个类别还你一个实数 —— 一个 logit。可以是负的、正的、极大的、极小的,没有任何约束。它们是一次矩阵乘法的输出,而矩阵乘法从没听说过概率这回事。

这里有五个。拖动滑块移动第三个分数,看这一行会怎么变 —— 以及它永远不会怎么变:

第三个分数 2.0

注意底部那个合计随着你拖动第三个分数而漂移。静止时是 1.8;把 时是 5.8。没有任何东西在把它往某处拉,也不可能有 —— 产生这些数的那一层里,根本没有归一化这一步。

要变成分布,需要每个值都为正、且总和钉在一。取指数买到了前一半,但要付代价 —— 抬高第三个分数,看第二行:

第三个分数 2.0

看这一行的形状塌得多快。分数为 2 时最大的指数是 7.39,总和 12.38;到 时是 408.42 里的 403.43,另外四根柱子已经贴到基线上了。取指数不只是把数变正:它把分数之差变成了比值,而这个比值每多一个单位的分数就乘以一个 e。

一次除法买到后一半。把每个指数除以它们的总和,这一行就变成了一个概率分布,不管原来的分数是什么:

第三个分数 2.0

这就是完整的 softmax:先取指数,再除以总和。Σp 在滑块的每一个位置都读作 1.0000,因为分母正是分子里那些数的和。这个不变量不是近似,也不是学出来的 —— 它就是算术,在没训练过的网络上同样成立。

在取指数前把分数除以一个常数,就得到一个控制结果锐利程度的旋钮。在分数纹丝不动的情况下,拖动温度走完整个范围:

温度 1.00

在 时,赢家读到小数点后四位是 1.0000,这张图在做 argmax;T = 4 时最大概率是 0.2905,五个几乎持平。分数一次都没动过。温度不是模型的属性 —— 它是读取模型的方式,在采样时才选定。

所以 logit 是一个分数,而概率是对一整组分数的读法。下面所有内容讲的都是这两句话之间的落差,因为跨越它的方式决定了一次训练跑出来的是一个数还是一个 NaN。

02

两种损失,两种形状

损失就是一个形状。选哪个形状,决定了模型错得离谱时梯度会怎么做。

平方误差是每个人最先遇到的损失,用在回归上它也确实是对的:它是高斯噪声下的极大似然损失,梯度是 2(ŷ − y),并且关于预测值是凸的。

沿着这只碗拖动标记,把损失读出来 ——目标值是 1 处那条虚线:

预测值 2.20. 左右拖动移动标记;方向键单步移动,Home 键回到初始状态。
预测值 2.20

注意,干活的全是这个形状。惩罚随偏差的平方增长,所以差 2 的预测代价是差 1 的四倍,而梯度随误差线性增长:你离得越远,被推回来的力气越大。

分类问的是另一个问题。那里模型给正确类别报出一个概率,我们要惩罚的是「自信地错」。把标记往零的方向拖:

概率 0.600. 左右拖动移动标记;方向键单步移动,Home 键回到初始状态。
概率 0.600

看左边那堵墙。模型笃定且正确时 −log p 是 0,五五开时是 0.69, p = 0.01 时是 4.61,而当 p → 0 它发散。这就是对 one-hot 目标的交叉熵:模型给出的那个答案的意外程度,它无上界是故意的。

把平方误差搬到同一个问题上,差别就不再细微了。两条曲线读的是同一个概率;把标记往左拖,拖到模型「自信地错」的地方:

概率 0.500. 左右拖动移动标记;方向键单步移动,Home 键回到初始状态。
概率 0.500

时交叉熵收 3.00,平方误差收 0.90。更糟的是,不管模型错到什么程度,平方误差都被 1 卡住,于是两者分歧最大的地方,恰恰是最要紧的地方。

不过训练用的不是损失值,而是梯度 —— 平方误差正是在这里彻底失效。把标记拖到一个错得离谱的分数上:

分数 0.0. 左右拖动移动标记;方向键单步移动,Home 键回到初始状态。
分数 0.0

因为平方误差是套在 sigmoid 外面的,它的梯度带着一个 σ(1 − σ) 因子,而这个因子恰好在模型错得最厉害的地方归零。在、目标为 1 时,交叉熵以 0.9975 的力推,平方误差只有 0.00246 —— 弱 405 倍。饱和的单元不再学习,而这次跑看起来就像已经收敛了。

这个抵消正是分类要用交叉熵的理由:链式法则的 σ′ 项被 −log p 求导的 1/p 项精确抵消,剩下的只有 p − y。第 06 节看它掉出来。

03

指数把浮点数用完了

exp 是整条流水线里长得最快的东西,而它落进去的那个格式是有限的。

训练里的每个数都住在 IEEE-754 浮点数里。binary32 最大能装到3.4028235e38,1.4012985e-45 以下什么也装不了。这两个常数就是 exp 会撞上的墙。

下面这把梯子一格一个数量级,从最小的次正规数一直到超过 float32 的最大有限值。拖动滑块,看e 的分数次幂往上爬:

分数 2

往上一格是十倍,所以这根柱子的高度是这个数的指数,而不是它的大小。 exp(2) 刚好在 1 上面一点; 已经冲出顶部,而越过上限之后所有值都是同一个值:inf。中间没有渐进的退化 —— exp(88.72) 是有限的 3.39e38,exp(88.73) 就是无穷。

混合精度训练把大多数张量搬到十六位,而 float16 只给指数留了五位, float32 给了八位。把分数往上推,看哪一条上限先到:

分数 8

float16 顶在 65,504,也就是 —— 分数是十一,不是八十九。这就是为什么在矩阵乘法跑半精度时,各家的 autocast 策略都把 softmax、log_softmax、cross_entropy 留在 float32 里:乘法在十六位下是安全的,指数不是。

现在把第 01 节那条流水线开进这堵墙。把最大的那个分数抬过 88.72,三行一起读:

最大的分数 2

因为分子和分母一起溢出,赢的那个类别算的是 inf ÷ inf,得到 NaN;其余每个类别算的是有限数 ÷ inf,得到一个干干净净、看着很合理的 0。损失是 NaN。一次反向传播之后,碰过它的每个参数也都是 NaN —— 而且没有任何东西抛异常。

这才是真正要命的失败模式,因为它是无声的。没有异常、没有告警、没有半截结果:训练照跑,损失接下来一路打印 nan,存下来的 checkpoint 一文不值。挡在你和它之间的,就是下一节那两行。

04

把最大值减掉

一个恒等式让整个问题消失,代价是在这一行上多扫一遍。

对任意标量 c 都有 softmax(z) = softmax(z − c)。证明只有一行:分子分母同乘 e−c 什么也没变,而 ez−c = eze−c。这件事值得亲眼看一遍,而不是信一遍。

把五个分数整体平移同一个量,再看它下面那行概率:

平移量 0

注意下面那一行纹丝不动。上面每根柱子都挪了,而第三个概率稳稳停在 0.5967。 softmax 读的只是分数之间的差;绝对高度是它根本拿不到的信息。

所以 c 归我们挑,而有一个选法是特别的。把平移量拖到最大的那个分数上,看这些指数:

平移量 0

当 c = 最大值时,平移后最大的分数恰好是 0,于是它的指数恰好是 1,其余每一个都落在 (0, 1] 里。没有东西能溢出,因为 exp 拿到手的最大参数就是零 —— 而分母至少是 1,所以也没有东西能除以零。

把两条路线放到同一把梯子上。把原始分数拉到滑块的尽头,看哪一根柱子还留在梯级上:

最大的分数 2

原始那根柱子在 89 处冲出梯顶,再也没回来。平移后那根在任何分数下都读作 1,因为平移之后最大的指数永远是 e0。这份保证的代价是多扫一遍去找最大值:这一行扫三遍而不是两遍。

这就是为什么你永远不该照定义写 softmax。exp(z) / exp(z).sum() 是正确的公式和坏掉的程序;减掉最大值的那行是同一个公式和不可能失败的程序。各家框架的 softmax 都是后者。

05

log-sum-exp,以及那个永远不发生的 log

平移修好了 softmax,却修不好 log(softmax) —— 而损失要的正是这个 log。

交叉熵是 −log p,所以训练要的是对数概率。先算 softmax 再取 log 是一行代码,也是一个 bug:概率必须在浮点数里走完这个来回,而对自信的模型来说它走不完。

先看替代它的那个东西。log Σ exp 是一个软最大值 —— 最大的那个分数,加上一个只取决于差值的修正项。拖动标记,盯着这个修正项:

差值 0.00. 左右拖动移动标记;方向键单步移动,Home 键回到初始状态。
差值 0.00

差值为 0 时修正项是 ln 2 = 0.6931:两个分数相等,所以和是最大值的两倍。到差值 时它是 0.0003。 log-sum-exp 就是把拐角磨圆了的 max,而磨圆只发生在前两名接近的那一小段区域里。

再看这个 log 继承下来的失败。把两个分数的差值拉开,看第二名的概率沿着梯子往下走:

差值 20

低于 1.1754944e-38,float32 就离开了正规数范围,开始拿尾数位撑着走;低于 1.4012985e-45,它什么也不剩,直接返回 0。这件事发生在 处,而 104 不算什么 —— 训练好的语言模型在最高和最低 logit 之间放一百个 nat 是家常便饭。

恰好为零的概率,它的对数是 −inf,这正是朴素路线报出来的东西。下面两行算的是同一个量;把差值拖过 104:

差值 20

因为 z − log Σ exp 从来不构造那个概率,它也就不必表示它。在 处,它返回 −140.0,而 log(softmax) 返回 −inf —— −140 是一个再普通不过的 float32。减法是精确的;有损的那一步是指数,而我们把它跳过了。

log_softmax 不是 log(softmax(x)) 的顺手封装。它是另一套计算、另一套数值范围,这就是它作为独立算子存在的理由。

06

永远不要对原始分数取指数

这就是为什么各家框架的损失函数收的是 logits,而你若递给它概率,它会不声不响地毁掉这次训练。

从 logits 出发的交叉熵是 logΣexp(z) − z[y]:一次归约、一次减法,不构造概率,也不做除法。这个设计真正赚回本钱的地方,是它的梯度。

它对每个 logit 的导数是 p − y ——概率那一行减去one-hot 目标那一行。拖动正确类别的分数,用上面两行把第三行读出来:

分数 2.0

注意最下面那行确实是上面两行的逐分量之差:静止时目标那一列是 0.5967 − 1 = −0.4033,其余各列就是原本的概率。没有 σ′ 活下来,因为 log 带来的 1/p 和指数带来的 p 精确抵消了。不管分数是多少,每个分量都落在 [−1, 1] 里 —— 这就是交叉熵的梯度不会自己炸掉的原因。

如果你递给它的是概率那一行而不是分数那一行,它会再做一次 softmax —— 而且是不声不响地做,因为概率本身就是完全合法的浮点数。切换喂给损失的东西:

一次 softmax

看那条上限浮出来。第二次 softmax 拿到的输入已经被压进 [0, 1],所以它能看到的最大差值只有 1,赢家被卡在 e / (e + n − 1),五个类别时是 0.4046。模型照样训练,损失照样下降;它只是不可能低于这里的 0.9048,对 50,257 词表则不可能低于 9.82 —— 而均匀乱猜是 10.82。

这笔交易的另一半是:融合形式压根没有溢出可言。损失不可能依赖于平移量,所以它必须是平的;把最大的分数拖过 88.72,看哪条路线同意这一点:

最大的分数 20. 左右拖动移动标记;方向键单步移动,Home 键回到初始状态。
最大的分数 20

因为 logΣexp 在内部就平移过了,融合路线在 处依然读作 0.5163,而 −log(softmax) 从 88.73 起就是 NaN。在上限以下两者小数点后四位一致;在上限以上,一个是数,另一个不是。

上面还要再叠一条内存的账。 8 × 1,024 个位置、50,257 词表的概率张量,float32 下是 1.53 GiB,它的梯度又是 1.53 GiB —— 每一步都要分配、写入、释放一遍。融合算子两个都不构造。这就是这个 API 收 logits 的原因。

07

三条上限,五行代码

同一个表达式,是不是一个数由格式说了算。

整页讲的都是一个指数撞上一个有限格式。拖动分数,看哪些格式还装得下结果:

分数 11

float16 在 处就交待了,float32 在 88.72, float64 在 709.78 —— 整个冲出了这把梯子的顶。 bfloat16 和 float32 共用八位指数,所以同样能到 88.72,代价是精度从七位十进制降到约三位。

上面所有内容都装在五行里。第一行是唯一不属于教科书定义的那行,也正是它让后面四行安全:

m    = z.max()                  # the shift
lse  = m + log(exp(z - m).sum())
logp = z - lse                  # log_softmax
loss = lse - z[y]               # cross-entropy
grad = softmax(z) - onehot(y)   # p - y

把它们融合既是数值问题也是内存问题,而开销随词表增长。把词表从一个小分类器拖到一个语言模型,读这两个总量:

词表 1,000. 左右拖动移动标记;方向键单步移动,Home 键回到初始状态。
词表 1,000

在下,融合那一步握着 1.53 GiB 的 logits;不融合那一步握着 4.60 GiB,因为概率和它的梯度各是同一形状的又一份拷贝。两条轴都是对数轴,所以两条线是平行的:比值在任何词表下都是 3,动的只是绝对开销。

四种不声不响的失败。对原始分数取指数会返回NaN,而不是抛异常。对 softmax 的输出取 log,一旦概率下溢就返回 −inf。把概率喂给期待 logits 的损失,会训出一个系统性不自信的模型,全程不报错。而把其中任何一步挪进 float16 省显存,会在你一行代码都没改的情况下把上限从 88.72 压到 11.09。