数值稳定性:浮点数的漏的抽象 本节摘要:浮点(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 节。
阅读完本节,你应当能够:
数值稳定性不是理论关切,它是训练成功与静默失败的分水岭。每一个你将来调试的严肃 ML bug,最终都会归结到浮点。
三种典型故障:
inf,9002 步所有梯度 nan,训练死亡。inf,因为 exp(100) 超 float32 上限;每个 ML 框架都用一个两行技巧解决,你却不知道。浮点数由三部分组成:符号位、指数、尾数(有效数字)。
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.0000001 与 1.0000002,但区分不了 1.00000001 与 1.00000002,7 位之后全是舍入噪声。float16 仅约 3 位,最大值 65,504——而 logits、梯度、激活经常超过它,小得吓人。
bfloat16 是 Google 对 float16 范围问题的回答:它有与 float32 相同的 8 位指数(同样的范围,直到 3.4e38),但只有 7 位尾数(精度比 float16 还低)。对训练神经网络,范围比精度更重要,所以 bfloat16 通常胜出。
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) < epsilon 或 math.isclose()。
两个相近的浮点数相减,有效位相互抵消,剩下的舍入噪声被提升为前几位:
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 算法或先中心化;对数概率全程在对数空间运算。
溢出(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)) 若无技巧就是个雷区。
直接算 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 归一化、交叉熵、序列模型对数概率求和、高斯混合、变分推断。
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] ← 与不剪裁完全相同,但计算安全
概率完全一致,计算是安全的。这不是优化,是正确性的要求。
nan(非数)与 inf(无穷)会病毒式传播:梯度更新里一个 nan 让权重变 nan,继而让后续所有输出 nan,一步之内训练死亡。
inf 来源:exp(大正数)、除以零 1.0/0.0、float32 累加溢出。nan 来源:0.0/0.0、inf - inf、inf * 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。
解析梯度(来自反向传播)可能有 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()。
现代 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)开始,梯度溢出就减半,若干步不溢出就加倍。
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 的原因。
梯度爆炸发生在梯度经多层指数增长时(RNN、深网、Transformer 常见)。单个大梯度能一步毁掉所有权重。两种裁剪:
[-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 能产出足以毁掉数周训练的巨大梯度。
批归一化、层归一化、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)在所有激活相同时防除零;可学的 gamma、beta 让网络恢复任意所需尺度。这让全网络的值处于数值安全范围,既防前向溢出又防反向梯度爆炸。
完整源码见 phases/01-math-foundations/13-numerical-stability/code/numerical.py。
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]
def logsumexp_stable(values): c = max(values) return c + math.log(sum(math.exp(v - c) for v in values)) # 永不溢出/下溢
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
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
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/。
E[x²] − E[x]² 在 float32 下算 [1000000.0, 1000001.0, 1000002.0] 的方差,再用 Welford 算法算一遍,与真值 0.6667 比较误差。x 使 1.0 + x == 1.0(机器 epsilon),验证它等于 numpy.finfo(numpy.float32).eps。logsumexp_stable 测试三种情形:(a) 所有的值相等;(b) 一个值远大于其余;(c) 所有的值极负(-1000),验证朴素版失效处稳定版仍正确。y = Wx + b 及其解析反向,用 numerical_gradient 对 3×2 权重矩阵验证正确性。0.1 + 0.2 != 0.3:二进制无法精确表示 0.1;永远不要用 == 比较浮点,改用 abs(a-b) < eps。log(sum(exp(x))) = max(x) + log(sum(exp(x-max(x)))),减最大值消除溢出与下溢。max(logits) 后概率完全一致但永不溢出,log_softmax() 内部就是它。(f(x+h)-f(x-h))/(2h) 是 O(h²),相对误差 <1e-7 完美、>1e-3 有错。max_norm=1.0,是安全机制不是 hack。下一节,我们换一个视角看「大小」——范数与距离:L1/L2/余弦/马氏/Jaccard/编辑距离各自定义了什么叫「相似」,选错距离函数,下游一切都会塌方。