数值稳定性:浮点数的漏的抽象


文档摘要

数值稳定性:浮点数的漏的抽象 本节摘要:浮点(Floating Point)是一个会漏的抽象——它会在训练时咬你一口,而且你根本看不到它来。模型训了三小时损失变 NaN;或训完了却比论文低 2%,因为论文用 float32 而你用了没加 loss scaling 的 float16;或你手写交叉熵,logits 一过 100 就返回 ,因为 溢出了 float32。本节从 IEEE 754 讲起,揭穿 的真相与灾难性抵消(两个相近数相减丢掉有效位);给出 ML 里最关键的 log-sum-exp 技巧——减去最大值再做指数,从此 softmax 与交叉熵永不溢出;推导混合精度训练(float32 主权重 + float16 前向 + loss scaling 防梯度下溢);解释为何

数值稳定性:浮点数的漏的抽象

本节摘要:浮点(Floating Point)是一个会漏的抽象——它会在训练时咬你一口,而且你根本看不到它来。模型训了三小时损失变 NaN;或训完了却比论文低 2%,因为论文用 float32 而你用了没加 loss scaling 的 float16;或你手写交叉熵,logits 一过 100 就返回 inf,因为 exp(100) 溢出了 float32。本节从 IEEE 754 讲起,揭穿 0.1 + 0.2 != 0.3 的真相与灾难性抵消(两个相近数相减丢掉有效位);给出 ML 里最关键的 log-sum-exp 技巧——减去最大值再做指数,从此 softmax 与交叉熵永不溢出;推导混合精度训练(float32 主权重 + float16 前向 + loss scaling 防梯度下溢);解释为何 bfloat16 在训练中胜过 float16(同 float32 的范围、更少精度);并附上梯度检查(中心差分)、梯度裁剪、归一化层作为数值稳定器的完整工具箱。读完本节,NaN 不再是玄学。

对应原课程:Phase 01 · Lesson 13 · numerical-stability(原英文 phases/01-math-foundations/13-numerical-stability/docs/en.md)。前置:第 1~4 节。

学习目标

阅读完本节,你应当能够:

  1. 减最大值技巧实现数值稳定的 softmax 与 log-sum-exp。
  2. 识别浮点计算中的溢出、下溢、灾难性抵消
  3. 中心差分对照解析梯度与数值梯度。
  4. 解释为何 bfloat16 优于 float16,以及 loss scaling 如何防止梯度下溢。

一、问题与直觉

数值稳定性不是理论关切,它是训练成功与静默失败的分水岭。每一个你将来调试的严肃 ML bug,最终都会归结到浮点。

三种典型故障:

  • 突然 NaN:第 9000 步 logits 还正常,9001 步变 inf,9002 步所有梯度 nan,训练死亡。
  • 精度悄悄流失:架构、超参、数据都对,但你用了 float16 没加 scaling,32 位累积舍入误差悄悄吃掉了准确率。
  • exp 溢出:手写交叉熵,logits 过 100 就 inf,因为 exp(100) 超 float32 上限;每个 ML 框架都用一个两行技巧解决,你却不知道。

1.1 IEEE 754:计算机如何存实数

浮点数由三部分组成:符号位、指数、尾数(有效数字)。

float32 布局(共 32 位): [1 符号] [8 指数] [23 尾数] 值 = (-1)^sign * 2^(exponent - 127) * 1.mantissa

尾数决定精度(多少位有效数字),指数决定范围(数能多大或多小)。

格式 位 指数 尾数 十进制有效位 范围(约) float64 64 11 52 ~15-16 +/- 1.8e308 float32 32 8 23 ~7-8 +/- 3.4e38 float16 16 5 10 ~3-4 +/- 65,504 bfloat16 16 8 7 ~2-3 +/- 3.4e38

float32 给你约 7 位十进制精度——能区分 1.00000011.0000002,但区分不了 1.000000011.00000002,7 位之后全是舍入噪声。float16 仅约 3 位,最大值 65,504——而 logits、梯度、激活经常超过它,小得吓人。

bfloat16 是 Google 对 float16 范围问题的回答:它有与 float32 相同的 8 位指数(同样的范围,直到 3.4e38),但只有 7 位尾数(精度比 float16 还低)。对训练神经网络,范围比精度更重要,所以 bfloat16 通常胜出。

1.2 为什么 0.1 + 0.2 != 0.3

0.1 在二进制浮点里无法精确表示,它是无限循环小数:0.1 = 0.0001100110011...。float32 截断到 23 位尾数,存的是约 0.100000001490116;0.2 存成约 0.200000002980232;两者和是 0.300000004470348,不是 0.3。

>>> 0.1 + 0.2 0.30000000000000004 >>> 0.1 + 0.2 == 0.3 False

这对 ML 的影响:① if loss < threshold 可能给错答案;② 累加大量小值(几千步梯度更新)会偏离真和;③ 用 == 比较浮点的复现测试会失败。修复:永远不要用 == 比较浮点,改用 abs(a-b) < epsilonmath.isclose()

1.3 灾难性抵消(Catastrophic Cancellation)

两个相近的浮点数相减,有效位相互抵消,剩下的舍入噪声被提升为前几位:

a = 1.0000001(存为 1.00000011920929) b = 1.0000000(存为 1.00000000000000) 真差: 0.0000001 算得: 0.00000011920929 相对误差: 19.2% ← 一次减法丢掉近 20% 精度

ML 中这发生在:算均值很大数据的方差(E[x²] − E[x]²,当 E[x] 大时);减两个相近的对数概率;用太小的 epsilon 算有限差分梯度。修复:重排公式避开大数相减——方差用 Welford 算法或先中心化;对数概率全程在对数空间运算。

1.4 溢出与下溢

溢出(Overflow)是结果太大无法表示;下溢(Underflow)是太接近 0(小于最小可表示正数)。

float32 边界: 最大值: 3.4028235e+38 最小正(正规): 1.175e-38 最小正(非正规): 1.401e-45 溢出: > 3.4e38 -> inf 下溢: < 1.4e-45 -> 0.0 exp() 是 ML 溢出的主源: exp(88.7) = 3.40e+38 (勉强塞进 float32) exp(89.0) = inf (溢出) exp(-87.3) = 1.18e-38 (勉强高于下溢) exp(-104) = 0.0 (下溢归零) log() 反方向: log(0.0) = -inf log(-1.0) = nan log(1e-45) = -103.3 (没问题) log(1e-46) = -inf (输入先下溢为 0,再 log(0)=-inf)

ML 里 exp() 出现在 softmax、sigmoid、概率计算;log() 出现在交叉熵、对数似然、KL 散度。组合 log(exp(x)) 若无技巧就是个雷区。

1.5 log-sum-exp 技巧

直接算 log(sum(exp(x_i))) 数值上很危险:任一 x_i 大,exp(x_i) 溢出;全都极负,每个 exp(x_i) 下溢为 0,log(0)-inf

技巧:先减最大值再指数。

log(sum(exp(x_i))) = max(x) + log(sum(exp(x_i - max(x))))

为何有效:减 max(x) 后最大指数是 exp(0)=1,不可能溢出;和式中至少一项为 1,故和至少为 1,log(1)=0,不可能下溢到 -inf

证明:

log(sum(exp(x_i))) = log(sum(exp(x_i - c + c))) (加减 c) = log(sum(exp(x_i - c) * exp(c))) (exp(a+b)=exp(a)*exp(b)) = log(exp(c) * sum(exp(x_i - c))) (提出 exp(c)) = c + log(sum(exp(x_i - c))) (log(a*b)=log(a)+log(b))

c = max(x),溢出被消除。这一技巧在 ML 中无处不在:softmax 归一化、交叉熵、序列模型对数概率求和、高斯混合、变分推断。

1.6 softmax 为何必须减最大值

softmax 把 logits 转概率:softmax(x_i) = exp(x_i) / sum(exp(x_j))。不加技巧,logits [100, 101, 102]exp(100) = inf(float32)。

加技巧,减 max(x)=102:

exp(100-102)=exp(-2)=0.135 exp(101-102)=exp(-1)=0.368 exp(102-102)=exp(0) =1.000 和 = 1.503 softmax = [0.090, 0.245, 0.665] ← 与不剪裁完全相同,但计算安全

概率完全一致,计算是安全的。这不是优化,是正确性的要求。

1.7 NaN 与 Inf:检测与预防

nan(非数)与 inf(无穷)会病毒式传播:梯度更新里一个 nan 让权重变 nan,继而让后续所有输出 nan,一步之内训练死亡。

inf 来源:exp(大正数)、除以零 1.0/0.0、float32 累加溢出。
nan 来源:0.0/0.0inf - infinf * 0、负数开方、负数取 log、任何含已有 nan 的运算。

检测:math.isnan(x)math.isinf(x)math.isfinite(x)

预防:① 把 exp() 的输入 clamp 到 [-80, 80];② 分母加 epsilon x/(y+1e-8);③ log() 里加 epsilon log(x+1e-8);④ 用稳定实现(log-sum-exp、稳定 softmax);⑤ 梯度裁剪防权重爆炸;⑥ 调试时每轮前向后查 nan/inf

1.8 数值梯度检查

解析梯度(来自反向传播)可能有 bug。数值梯度检查用有限差分验证之。中心差分公式:

df/dx ≈ (f(x + h) - f(x - h)) / (2h) ← O(h²),远优于前向差分 (f(x+h)-f(x))/h 的 O(h)

选 h:太大则近似失真,太小则灾难性抵消毁掉结果,典型 h = 1e-5 ~ 1e-7。检查方式是算解析与数值梯度的相对差:

relative_error = |grad_解析 - grad_数值| / max(|grad_解析|, |grad_数值|, 1e-8)

经验阈值:<1e-7 完美、<1e-5 可接受、>1e-3 有问题、>1 完全错。实现新层或新损失时总要检查;PyTorch 提供 torch.autograd.gradcheck()

1.9 混合精度训练

现代 GPU 有专用硬件(Tensor Core)把 float16 矩阵乘算得比 float32 快 2~8 倍。混合精度训练:

1. 维护 float32 主权重副本 2. 前向用 float16(快) 3. 算 loss 用 float32(防溢出) 4. 反向用 float16(快) 5. 梯度放大回 float32 6. 更新 float32 主权重

纯 float16 的问题:梯度常极小(1e-8 或更小),float16 把 ~6e-8 以下都下溢为零,模型停止学习。修复是 loss scaling:① loss 乘大因子(如 1024);② 反向算 (loss×1024) 的梯度;③ 所有梯度放大 1024 倍(推到 float16 下溢之上);④ 更新前除回 1024;净效果:同样更新,但不下溢。动态 loss scaling 自动调节:从大值(65536)开始,梯度溢出就减半,若干步不溢出就加倍。

1.10 bfloat16 vs float16:为何训练用 bfloat16

float16: [1 符号] [5 指数] [10 尾数] 最大 ~65,504,精度更高 bfloat16: [1 符号] [8 指数] [7 尾数] 最大 ~3.4e38,精度更低

float16 精度更高(10 vs 7 位尾数)但范围有限;bfloat16 精度更低但范围同 float32。对训练:① 激活与 logits 训练尖峰常超 65,504,float16 溢出,bfloat16 不怕;② float16 必须配 loss scaling,bfloat16 因范围覆盖梯度量级通常不需要;③ bfloat16 是 float32 的简单截断(丢末 16 位尾数),转换无指数损失、几近无损。float16 适合推理(值有界、精度重要),bfloat16 适合训练(范围重要)。这就是 TPU 与现代 NVIDIA GPU(A100、H100)原生支持 bfloat16 的原因。

1.11 梯度裁剪

梯度爆炸发生在梯度经多层指数增长时(RNN、深网、Transformer 常见)。单个大梯度能一步毁掉所有权重。两种裁剪:

  • 按值裁剪:每个梯度元素独立 clamp 到 [-max_val, max_val],简单但可能改变梯度向量方向。
  • 按范数裁剪:缩放整个梯度向量使其范数不超过阈值:if ||grad|| > max_norm: grad = grad * (max_norm/||grad||)。保持方向——这是 torch.nn.utils.clip_grad_norm_() 的做法,是标准选择。典型:Transformer max_norm=1.0,RL 0.5,简单网络 5.0

梯度裁剪不是 hack,是安全机制——没有它,一个异常 batch 能产出足以毁掉数周训练的巨大梯度。

1.12 归一化层作为数值稳定器

批归一化、层归一化、RMS 归一化通常被视为助收敛的正则化器,但它们也是数值稳定器。无归一化时激活会逐层指数膨胀:Layer 1 ∈ [0,1]Layer 5 ∈ [0,100]Layer 10 ∈ [0,10000]Layer 50 ∈ [0, inf]

归一化每层重定心、重缩放:

LayerNorm(x) = (x - mean(x)) / (std(x) + epsilon) * gamma + beta

epsilon(典型 1e-5)在所有激活相同时防除零;可学的 gammabeta 让网络恢复任意所需尺度。这让全网络的值处于数值安全范围,既防前向溢出又防反向梯度爆炸。

二、从零实现

完整源码见 phases/01-math-foundations/13-numerical-stability/code/numerical.py

2.1 朴素 vs 稳定 softmax

def softmax_naive(logits): exps = [math.exp(z) for z in logits] total = sum(exps) return [e / total for e in exps] # logits 含 100 就全 nan def softmax_stable(logits): max_logit = max(logits) exps = [math.exp(z - max_logit) for z in logits] # 减最大值 total = sum(exps) return [e / total for e in exps]

2.2 稳定 log-sum-exp

def logsumexp_stable(values): c = max(values) return c + math.log(sum(math.exp(v - c) for v in values)) # 永不溢出/下溢

2.3 稳定交叉熵

def cross_entropy_stable(true_class, logits): max_logit = max(logits) shifted = [z - max_logit for z in logits] log_sum_exp = math.log(sum(math.exp(s) for s in shifted)) log_prob = shifted[true_class] - log_sum_exp # = log softmax return -log_prob

2.4 梯度检查(中心差分)

def numerical_gradient(f, x, h=1e-5): grad = [] for i in range(len(x)): xp = x[:]; xm = x[:] xp[i] += h; xm[i] -= h grad.append((f(xp) - f(xm)) / (2 * h)) return grad def check_gradient(analytical, numerical, tol=1e-5): for i, (a, n) in enumerate(zip(analytical, numerical)): denom = max(abs(a), abs(n), 1e-8) rel = abs(a - n) / denom print(f" param {i}: 解析={a:.8f} 数值={n:.8f} 相对误差={rel:.2e} " f"[{'OK' if rel < tol else 'FAIL'}]")

三、框架对比

混合精度模拟

def simulate_bfloat16(x): packed = struct.pack('f', x) as_int = int.from_bytes(packed, 'little') truncated = as_int & 0xFFFF0000 # 截断末 16 位尾数 return struct.unpack('f', truncated.to_bytes(4, 'little'))[0]

按范数裁剪梯度

def clip_by_norm(gradients, max_norm): total_norm = math.sqrt(sum(g**2 for g in gradients)) if total_norm > max_norm: scale = max_norm / total_norm return [g * scale for g in gradients] # 保持方向 return gradients

NaN/Inf 检测

def check_tensor(name, values): has_nan = any(math.isnan(v) for v in values) has_inf = any(math.isinf(v) for v in values) if has_nan or has_inf: print(f"WARNING {name}: nan={has_nan} inf={has_inf}") return False return True

💡 永远用 torch.nn.functional.log_softmax() 而非手写 log(softmax())——前者内部实现 log-sum-exp,数值安全;后者先 softmax 可能下溢为 0 再 log 得 -inf

四、可复用产物

  • code/numerical.py:稳定 softmax、log-sum-exp、交叉熵、梯度检查、混合精度模拟的完整实现。
  • outputs/prompt-numerical-debugger.md:一份诊断 NaN/Inf 与数值问题的提示,贴进 AI 助手即可定位根因。

这些稳定实现在第 3 章搭训练循环、第 4 章实现注意力时会反复用到。源码见 phases/01-math-foundations/13-numerical-stability/code/

五、练习

  1. (Easy) 灾难性抵消:用朴素公式 E[x²] − E[x]² 在 float32 下算 [1000000.0, 1000001.0, 1000002.0] 的方差,再用 Welford 算法算一遍,与真值 0.6667 比较误差。
  2. (Medium) 精度猎手:找最小的正 float32 x 使 1.0 + x == 1.0(机器 epsilon),验证它等于 numpy.finfo(numpy.float32).eps
  3. (Medium) log-sum-exp 边界:用你的 logsumexp_stable 测试三种情形:(a) 所有的值相等;(b) 一个值远大于其余;(c) 所有的值极负(-1000),验证朴素版失效处稳定版仍正确。
  4. (Medium) 梯度检查神经网络层:实现单层 y = Wx + b 及其解析反向,用 numerical_gradient 对 3×2 权重矩阵验证正确性。
  5. (Hard) loss scaling 实验:模拟 float16 训练——生成 [1e-9, 1e-3] 范围的随机梯度,转 float16,测多少比例变零;再用 loss scaling(乘 1024)、转 float16、缩回,重测零比例。

本节要点回顾

  1. 浮点是漏的抽象:IEEE 754 用符号/指数/尾数存实数,float32 约 7 位精度,float16 仅 3 位且最大 65,504。
  2. 0.1 + 0.2 != 0.3:二进制无法精确表示 0.1;永远不要用 == 比较浮点,改用 abs(a-b) < eps
  3. 灾难性抵消:两个相近数相减丢有效位;重排公式(方差用 Welford、对数全程对数空间)规避。
  4. log-sum-exp 技巧:log(sum(exp(x))) = max(x) + log(sum(exp(x-max(x)))),减最大值消除溢出与下溢。
  5. 稳定 softmax 是正确性要求:减 max(logits) 后概率完全一致但永不溢出,log_softmax() 内部就是它。
  6. 梯度检查用中心差分:(f(x+h)-f(x-h))/(2h) 是 O(h²),相对误差 <1e-7 完美、>1e-3 有错。
  7. 混合精度 = float32 主权重 + float16 前后向 + loss scaling:scaling 放大梯度防 float16 下溢,更新前再缩回。
  8. bfloat16 训练胜过 float16:同 float32 的 8 位指数(范围到 3.4e38),通常不需 loss scaling;float16 适合推理。
  9. 梯度裁剪按范数:保持方向地缩放梯度向量,Transformer max_norm=1.0,是安全机制不是 hack。
  10. 归一化层是数值稳定器:每层重定心重缩放,防激活指数膨胀导致的前向溢出与反向梯度爆炸。

下一节,我们换一个视角看「大小」——范数与距离:L1/L2/余弦/马氏/Jaccard/编辑距离各自定义了什么叫「相似」,选错距离函数,下游一切都会塌方。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U