NUMERICAL STABILITY LAB

Softmax 为什么要先减最大值?

把同一组 logits 送进两条计算路径。概率的数学结果本该一样,但浮点数的表示范围会让“直接取指数”先出问题。

pₖ = exp(zₖ − max(z)) / Σ exp(zⱼ − max(z)) ∂L / ∂z = p − y

先把 logits 摆出来

输入范围是 −10000 到 10000。logit 不是概率,它可以为负,也不需要相加为 1。

改动任意数值后,两条路径会同步重算。

同一组数,两种算法

左边直接计算 ez;右边先用每个 logit 减去当前最大值。概率条的长度与 p 成正比。

直接指数计算

正常

pₖ = exp(zₖ) / Σ exp(zⱼ)

类别 A
类别 B
类别 C
类别 D
概率和 Σp
真实类交叉熵

减最大值后计算

稳定

pₖ = exp(zₖ − m) / Σ exp(zⱼ − m), m = max(z)

类别 A
类别 B
类别 C
类别 D
概率和 Σp
真实类交叉熵

再看输出层梯度 p − y

真实类别的 y 等于 1,其余类别等于 0。负值表示该 logit 需要被推高,正值表示当前概率会将对应 logit 往下压。

p − y = [—]
梯度和:—
类别 A
p = —,y = —
类别 B
p = —,y = —
类 C
p = —,y = —
类别 D
p = —,y = —