所属方向:分布式训练 | 难度:进阶 | 前置知识:无
关键词:混合精度、FP16、BF16、FP8、loss scaling、溢出、数值稳定
同一份训练代码,把 BF16 改成 FP16,吞吐可能更高,也可能几十步后 loss 直接变成 NaN。另一个常见现场是:模型看起来正常收敛,但某些参数长期没有变化;检查梯度才发现,一些绝对值不超过 $2^{-25}\approx2.98\times10^{-8}$ 的梯度在 round-to-nearest-even 写回 FP16 时已经变成了 0。
这不是“半精度不准”一句话能解释的。训练里的风险至少有三种:数太大,超过格式的最大有限值而溢出成 inf;数太小,舍入后落到 0(或硬件将次正规数 flush-to-zero);数虽然在范围内,却因为有效位太少,在加法里被大数吞掉。三者的修复手段不同:loss scaling 能救小梯度,却救不了激活溢出;换 BF16 能扩展动态范围,却不会自动改善尾数精度;把归约留在 FP32 能减小累加误差,却不代表所有输入和权重都要回到 FP32。
为什么还要承担这些麻烦?因为 Transformer 的大头通常是矩阵乘。现代加速器对低精度矩阵乘提供更高吞吐,权重和激活每元素从 4 字节降到 2 字节也能减小显存与带宽压力。Mixed Precision Training 给出的核心配方是:低精度做大部分前后向,保留 FP32 主权重进行更新,并对 loss 缩放以防小梯度消失;其模型显存可接近减半,同时保持精度。
“混合”的重点不是选一个统一 dtype,而是给不同数值角色分工。大 GEMM 适合低精度输入,softmax、归一化、loss 和长归约通常需要更大的范围或更高精度,优化器状态又有自己的要求。AMP 的价值正是把逐算子的 dtype 路由表和梯度缩放流程标准化。
一个数字能说明 FP16 与 BF16 的性格差异:FP16 最大有限值只有 65504,而 BF16 与 FP32 一样有 8 位指数,最大值约 $3.39\times10^{38}$。反过来,FP16 有 10 位小数尾数,1 附近的间隔约 $9.77\times10^{-4}$;BF16 只有 7 位尾数,间隔约 $7.81\times10^{-3}$。BF16 更不容易炸,FP16 在可表示范围内更细;“更稳定”与“更精确”不是同一维度。
三句话先建立框架:
如果只记一个检查顺序:先问异常发生在前向还是反向,再区分 overflow、underflow 与 rounding,最后才决定换 dtype、调 scale,还是把局部算子提升到 FP32。
AMP 中常见的四种角色也要分开:参数存储 dtype、算子输入 dtype、乘法累加 dtype、优化器状态 dtype。日志打印“BF16 training”并不能证明四者都是 BF16。很多 GEMM 是 BF16 输入、FP32 累加;Adam 的一阶矩和二阶矩仍为 FP32;某些框架还会保存 FP32 主权重。
对非零正规数,一个二进制浮点数可写成(次正规数没有隐含的前导 1):
$$x=(-1)^s(1.f)_2\,2^{e-\mathrm{bias}}$$
$s$ 是符号位,$e$ 是指数域,$f$ 是小数域。指数位数决定动态范围,尾数位数决定相邻可表示数的间距。FP16 是 1 位符号、5 位指数、10 位小数;BF16 是 1、8、7;FP32 是 1、8、23。于是:
机器 epsilon 是 1 与下一个可表示数的间隔;下面列的是 epsilon。在 round-to-nearest 模式下,通常定义的 unit roundoff 为这些值的一半:
$$\epsilon_{\mathrm{FP16}}=2^{-10},\qquad \epsilon_{\mathrm{BF16}}=2^{-7},\qquad \epsilon_{\mathrm{FP32}}=2^{-23}$$
因此 $1+10^{-4}$ 写入 FP16 或 BF16 都可能仍是 1。更危险的是参数更新:若权重 $w=1$,学习率乘梯度只有 $10^{-5}$,直接在低精度权重上做 $w-\eta g$,变化会被舍掉。FP32 主权重就是在高精度副本上累计细小更新,再把结果舍入给低精度前后向。
在 round-to-nearest-even 且保留次正规数时,绝对值不超过 $2^{-25}$ 的 FP16 输入舍入为零;介于这个阈值和 $2^{-24}$ 之间的数可能舍入为最小次正规数,而非一律为零。支持 flush-to-zero 的执行路径还需另查硬件与算子规则。
设损失为 $L(\theta)$,参数为 $\theta$,真实梯度为 $g=\nabla_\theta L$。把 loss 乘常数 $S$ 后反向:
$$g_s=\nabla_\theta(SL)=S\nabla_\theta L=Sg$$
只要在优化器读取梯度前除以 $S$,就恢复原梯度:
$$g=\frac{g_s}{S}$$
关键是量化顺序。若 $g=10^{-8}$,先写入 FP16 会变 0,之后再乘任何数都救不回来;若先由链式法则得到 $Sg=1.024\times10^{-5}$,它能被 FP16 表示,写回后再用 FP32 除以 1024,就能保留接近 $10^{-8}$ 的非零值。
静态 scale 要人工选 $S$。太小救不了下溢,太大又会让大梯度超过 65504。动态 GradScaler 维护随训练变化的 $S_t$:
$$S_{t+1}=\begin{cases}\beta S_t,&g_s\text{ 含 inf/NaN}\\ \gamma S_t,&\text{连续 }K\text{ 步有限}\\ S_t,&\text{其他情况}\end{cases}$$
$\beta<1$ 是回退因子,$\gamma>1$ 是增长因子,$K$ 是增长间隔。出现非有限梯度时必须跳过 optimizer step,否则坏更新会污染权重和 Adam 状态。PyTorch 文档还提醒:scale 不保证始终大于 1,BF16 预训练模型强转 FP16 时可能因数值过大而一路回退。
直接计算 softmax:
$$p_i=\frac{e^{x_i}}{\sum_j e^{x_j}}$$
若 $x_i$ 很大,指数会溢出。稳定写法先减最大值 $m=\max_jx_j$:
$$p_i=\frac{e^{x_i-m}}{\sum_j e^{x_j-m}}$$
这样最大指数为 1,不改变结果,却压下溢出风险。LayerNorm 方差也不宜用 $E[x^2]-E[x]^2$ 的低精度朴素形式:两个接近的大数相减会灾难性消减。长向量求和误差还会随项数累积,因此 AMP 通常让归约在 FP32 完成。
还要区分“算子输出 dtype”和“内部累加 dtype”。矩阵乘输入、输出可以是 BF16,但硬件乘积常进入 FP32 accumulator,最终再舍入回 BF16。若自定义 kernel 把 accumulator 也降成 16 位,它不再与常规 AMP 拥有同样数值性质。
FP8 常见 E4M3 与 E5M2。FP8 Formats for Deep Learning 用 E4M3 提供较多有效位,用 E5M2 提供较大范围。原始张量很难恰好落进狭窄范围,通常给每个张量或块维护尺度 $a$:
$$q=Q_{\mathrm{FP8}}(x/a),\qquad x_{\mathrm{deq}}=a\,q$$
尺度可由绝对最大值 amax 与格式上限估计。若每步都同步全局 amax,又会引入开销;工业实现会用历史窗口、延迟缩放或块缩放。FP8 的正确性同时依赖格式、粒度、尺度更新、异常值分布和高精度累加,不能机械替换 dtype 字符串。

直接量化的相对误差。小数值端 FP16 可能舍入为零(相对误差 1),超过最大有限值后溢出;BF16 的范围更宽,但正规数的尾数精度更低。精确可表示点的误差为零,图中作绘图下限处理。
两个脚本都只依赖 Python 标准库。float_formats.py 用 IEEE 二进制打包模拟 FP16,并先按舍入位加偏置,再截去 FP32 低 16 位,实现 BF16 round-to-nearest-even。真实输出:
format exp frac min_normal max_finite epsilon_at_1
FP16 5 10 6.1035e-05 6.5504e+04 9.7656e-04
BF16 8 7 1.1755e-38 3.3895e+38 7.8125e-03
FP32 8 23 1.1755e-38 3.4028e+38 1.1921e-07
value -> FP16 | BF16
1.0001 -> 1 | 1
1e-05 -> 1.00136e-05 | 1.00136e-05
100000 -> inf | 99840
100000 在 FP16 变成 inf,在 BF16 仍有限;1.0001 在两种 16 位格式里都舍入成 1。动态范围与精度的区别由此可见。
loss_scaling_sim.py 把五个小梯度直接量化,再先乘 1024、量化、最后除回去:
shape=(5,), loss_scale=1024
gradient direct_fp16 scaled_fp16/unscaled
1.000e-08 0.000e+00 1.001e-08
3.000e-08 5.960e-08 2.998e-08
1.000e-07 1.192e-07 1.000e-07
1.000e-06 1.013e-06 1.000e-06
1.000e-05 1.001e-05 9.999e-06
nonzero: direct=4/5, scaled=5/5
$10^{-8}$ 直接写 FP16 已为 0,缩放路径却保留为 $1.001\times10^{-8}$。这不是凭空增加精度:量化误差仍在,只是把数搬进可表示区间。
真实 PyTorch CUDA 训练的核心顺序:
scaler = torch.amp.GradScaler("cuda")
for inputs, targets in loader:
optimizer.zero_grad(set_to_none=True)
with torch.autocast("cuda", dtype=torch.float16):
loss = loss_fn(model(inputs), targets)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
梯度裁剪必须在 unscale 之后,若对放大后的梯度仍用阈值 C 裁剪,反缩放后的有效阈值会变成 C/S。若用 BF16,通常保留 autocast 而不用 GradScaler;是否省略仍应由实际梯度分布验证。
以 2026-09 的实现为准,PyTorch 的 torch/amp/grad_scaler.py 中 GradScaler.scale 把 loss 乘当前尺度;unscale 按 device 与 dtype 分组检查并反缩放梯度;step 只在没有 inf/NaN 时调用优化器;update 再根据各 optimizer 的 found-inf 状态调整尺度。它还处理稀疏梯度、多设备、多优化器、checkpoint state dict 和惰性初始化,这些都是最小模拟没有覆盖的工程边界。
官方 AMP 文档 当前推荐统一使用 torch.autocast 与 torch.amp.GradScaler;旧的 torch.cuda.amp 接口已标记弃用。autocast 只包前向和 loss,backward 放在上下文外。不要启用 autocast 后再全局 model.half,否则会绕开逐算子的安全策略。
算子策略不是“白名单里全降精度”。线性层、卷积通常进入 lower-precision;softmax、部分 loss 和归约倾向 FP32;多输入算子可能提升到最宽输入类型。显式 dtype、原地算子或 out 版本可能不参与 autocast,自定义代码要检查实际 dtype 流。
FP8 方面,NVIDIA Transformer Engine 提供 E4M3/E5M2 recipe、amax 历史与 delayed scaling,并在支持硬件上让 Transformer 层进入 FP8 路径。它仍以 BF16/FP16 保存部分张量并用高精度累加;分布式训练还可能跨 rank 归约 amax。“开启 FP8”是引入量化运行时,不是把参数永久存成单一 8 位格式。
生产监控至少记录当前 scale、增长和回退次数、跳过的 step 数、梯度范数、参数与激活 amax、NaN 首次出现的层、关键算子 dtype。只看总 loss 会把静默下溢藏很久。
混合精度通常节省权重、梯度和激活带宽,但不保证总显存恰好减半。Adam 的 FP32 主权重与两个 moment 仍可能占 12 字节/参数;部分算子保存 FP32 中间量;通信 buffer、cast buffer、workspace 与碎片也增加峰值。
吞吐提升也有条件。矩阵尺寸太小、CPU 或数据加载受限、频繁转换、未使用低精度加速单元,都会让 AMP 几乎不加速。小 batch 下,kernel launch 与 cast 成本甚至抵消收益。正确基准应同时报告 tokens/s、峰值显存、收敛曲线和最终指标。
以下情况尤其要谨慎:
调试时先用 FP32 建立可复现基线,再开 BF16,然后才是 FP16/FP8。每步只改一个变量,并比较前若干 step 的 loss、梯度范数和权重更新。数值问题最怕同时调整学习率、并行度、batch 与 dtype。
以 AdamW 为例,一次更新至少涉及前向激活、反向梯度、参数、副本和两个动量。低精度 GEMM 只覆盖其中一部分。设低精度参数为 $\theta_{16}$,FP32 主参数为 $\theta_{32}$,反缩放后的梯度为 $g_{32}$,则更新链条可以写成:
$$\theta_{32}^{t+1}=\operatorname{AdamW}(\theta_{32}^{t},g_{32}^{t},m_t,v_t),\qquad \theta_{16}^{t+1}=Q_{16}(\theta_{32}^{t+1})$$
$m_t,v_t$ 是 FP32 一阶、二阶矩,$Q_{16}$ 表示舍入到前后向 dtype。若框架没有主参数副本,而直接更新 BF16 参数,许多小于当前参数 ULP 的更新会消失。是否保留主权重是优化器实现细节,不能仅从模型参数 dtype 推断。
梯度累积又多一层顺序。假设一个 optimizer step 包含 $A$ 个 microbatch,它们必须使用同一 scale $S$,先把已缩放梯度累加:
$$g_s=\sum_{a=1}^{A}Sg_a=S\sum_{a=1}^{A}g_a$$
等所有 microbatch 完成后只反缩放一次。若中途 update scale,累加 buffer 中不同 microbatch 带有不同倍数,最后无法用一个 $S$ 恢复。若某个 microbatch 产生 inf,这一整个 optimizer step 都应跳过,而不是只丢掉坏 microbatch,否则有效 batch 和采样权重已变化。
数据并行中,每个 rank 只看到本地样本。rank 3 出现 inf 而其他 rank 有限时,所有副本仍必须对“是否更新”达成一致;否则 rank 3 跳步、其余 rank 更新,参数立即分叉。工业 GradScaler 会把 found-inf 状态跨相关设备和优化器汇总,分布式封装还需确保所有 DP rank 做相同决策。
梯度裁剪也有全局语义。ZeRO/FSDP 把梯度分片后,每卡只能算局部平方和:
$$\|g\|_2=\sqrt{\sum_{r=0}^{P-1}\sum_{i\in\mathcal{S}_r}g_i^2}$$
必须 all-reduce 局部平方和才能得到全局范数。正确顺序通常是:完成梯度同步,反缩放,检查非有限值,计算全局 norm,裁剪,执行 optimizer step,最后更新 scale。任一步调换都可能让“裁剪阈值 1.0”失去原含义。
混合精度还会改变 collective 的数值误差。BF16 梯度 all-reduce 比 FP32 少一半字节,但不同归约树改变加法顺序;卡数越多、梯度尺度跨度越大,末位差异越明显。若训练对归约误差敏感,可使用 FP32 梯度通信或分块高精度累加,但要接受带宽代价。排查时应分开比较“本地梯度生成精度”与“跨 rank 归约精度”。
选 FP16 还是 BF16,不应只看硬件宣传峰值。可按四步判断:
第一,确认硬件原生路径。某些设备对 BF16 与 FP16 吞吐相同,某些旧设备没有 BF16 Tensor Core;软件模拟 BF16 可能更慢。第二,检查模型来源。BF16 预训练 checkpoint 的激活范围未必适合 FP16,直接转换容易前向溢出。第三,检查数值敏感区域。softmax、归一化、概率 loss、长归约优先保留 FP32。第四,跑短程 A/B,比较吞吐、峰值、跳步率和收敛,而非只确认“能启动”。
可以把训练方案分成四档:
所谓“纯 BF16”或“纯 FP8”常是营销简写。softmax、norm statistic、optimizer state、某些 residual 累加和 master gradient 仍可能更高精度。真正有意义的配置描述,应列出 input、weight、output、accumulator 和 optimizer 五个维度。
排查顺序应尽量靠近异常源。先在每个模块前后记录有限值比例、绝对最大值、绝对非零最小值与 dtype;找到第一个从有限变非有限的边界。若 attention score 在 softmax 前已经 inf,检查 Q/K 范数、缩放因子与位置编码;若 loss 有限而 backward 首先 inf,检查导数奇点、自定义 backward 和 loss scale;若梯度有限但 step 后参数异常,检查 optimizer state、权重衰减和反缩放顺序。
静默下溢比 NaN 更难发现。建议记录零梯度比例,并按参数组看更新比率:
$$r_{\mathrm{update}}=\frac{\|\Delta\theta\|_2}{\|\theta\|_2+\varepsilon}$$
若某些层长期为 0,而 FP32 基线非零,说明低精度存储或 scale 窗口吞掉了更新。若所有层的比例突然增大,则更像学习率、loss normalization 或跳步恢复后的 scale 问题。
不要看到 NaN 就立即降低学习率。学习率过大确实可能导致发散,但 dtype overflow、错误 mask 产生全负无穷、空 batch 的除零、坏数据和通信错误都可能表现为同一个 NaN。先定位首次非有限张量,才能选择对应修复。
恢复训练时必须保存 GradScaler 状态。若只恢复模型和 optimizer,却把 scale 重置为很大初值,前几步可能反复溢出并被跳过;重置太小则产生额外下溢。保存点最好位于完整 optimizer step 之后,避免记录到一半累积的混合状态。
activation checkpoint 会在 backward 重跑前向。重算必须进入与原前向一致的 autocast 上下文,否则保存路径是 BF16、重算路径变 FP32 或 FP16,梯度对不上。torch.compile、CUDA Graph 和 fused optimizer 还可能改变 autocast 边界或引入 CPU-GPU 同步,升级版本后要重新做数值与性能回归。
自定义 autograd Function 不能假定输入永远 FP32。前向若内部要求 FP32,应显式关闭 autocast 并转换输入;反向要返回与调用约定兼容的梯度。写自定义 CUDA kernel 时,明确 accumulator 类型、饱和或 inf 行为、次正规数处理以及随机舍入。一个算子“支持 half”只代表能执行,不代表适合训练。
最后,评估数值稳定不能只跑几十步。动态 scale 的增长周期可能是数千步,某些罕见 batch 才触发极端值。至少保留一次覆盖学习率峰值、warmup 结束和验证阶段的长程对照,并把跳步数作为训练产物记录;否则断点续训后出现指标差异,很难追溯是数据还是精度状态造成。
AMP 基准至少要有三次 warmup 与多次稳定迭代,计时前后做设备同步;否则异步 kernel 会让 CPU 提交时间冒充 GPU 执行时间。显存要同时报告 allocated 与 reserved 峰值,前者是活跃张量,后者包含分配器缓存。对比方案必须使用同一 batch、相同梯度累积与相同 checkpoint 策略。
吞吐最好报告有效 token/s,而不是 batch/s。变长序列下,一个 batch 的 padding 比例不同,batch/s 会误导。质量侧至少比较训练 loss、验证指标、梯度范数和被跳过 step 数。若 AMP 每秒更快却需要更多 step 才到相同验证指标,最终 time-to-quality 未必更优。
对 FP8 还应增加量化覆盖率:多少 GEMM 真正走 FP8,多少因 shape、算子或 recipe 回退到 BF16;记录每层 amax、饱和比例和 scale 更新。仅看配置显示 FP8 无法证明加速路径被命中。kernel trace 能确认实际指令与 cast 开销。
建立门禁时,不必要求低精度与 FP32 每步完全相同。可先规定前 100 步 loss 相对误差、梯度 cosine、最终指标容差与最大跳步率,再用多个随机种子评估。数值差异是浮点并行的正常现象,持续偏向、层级爆炸或指标显著退化才是故障信号。
上线前逐项回答:硬件是否原生支持目标 dtype;矩阵尺寸是否对齐加速 tile;敏感算子是否保持 FP32;FP16 是否启用动态 scaling;裁剪是否在 unscale 后;所有 DP rank 是否共享跳步决策;梯度累积期间 scale 是否不变;checkpoint 是否保存 scaler;自定义算子是否声明 autocast 与 accumulator;监控是否能定位第一处非有限值。
然后做三个故障注入。人为把 scale 调得极大,确认系统发现 inf、跳过更新并回退;给输入加入一个幅值异常样本,确认前向监控能定位层;从 checkpoint 恢复,确认 scale、optimizer 与随机数状态连续。没有做过故障注入的告警,往往只在真正长跑失败后才发现无效。
最后保留 FP32 或 BF16 安全开关。生产训练发生异常时,能在不改数据顺序和并行布局的条件下提升精度复现,定位效率远高于临时改一堆超参。安全路径不一定长期运行,但必须定期测试,避免代码演进后早已失效。
这些流程看似比设置一个 autocast 开关繁琐,却能把“偶尔 NaN”“换卡就掉点”“断点后不收敛”变成有指标、有复现实验、有回退方案的普通工程问题。
一次训练 A 比 B 的 loss 高,并不能证明精度方案更差。数据顺序、dropout、并行归约顺序和非确定 kernel 都会制造波动。公平实验要固定数据索引与初始化,尽可能使用确定算法,并运行多个随机种子。先比较同一步同一层的输出、梯度与更新,再比较长程最终指标。
可用 FP32 输出 $y_{32}$ 作为局部参照,计算相对误差与余弦相似度:
$$e_{\mathrm{rel}}=\frac{\|y_{\mathrm{low}}-y_{32}\|_2}{\|y_{32}\|_2+\varepsilon}$$
单层误差略大未必影响训练,但若误差沿深度单调放大,往往说明 residual 累加、归一化或某个敏感算子精度不足。把统计按层绘制比只比较最终 logits 更容易定位。
还应区分可复现性与正确性。集合通信改变加法顺序后,bitwise 结果可能不同,但两条训练曲线仍落在相同统计分布。反之,两次运行逐位一致也可能稳定地实现了错误缩放。验收既要有小规模数学 oracle,也要有多种子任务指标。
混合精度上线后,持续监控跳步比例和 scale 分布。数据配方变化、序列变长、加入新 loss 或更换初始化,都可能改变数值范围;一次验证通过不意味着未来配置永久安全。把 precision 当成模型配置的一部分进行版本化,才能复现每次训练。
一个实用原则是把所有隐式转换显式化到观测层:训练启动时抽样打印关键模块的参数、输入、输出与归约 dtype,同时保存硬件、驱动、框架版本。低精度 kernel 与 autocast 策略会随版本更新,旧实验结论不能无条件外推。版本升级后的第一件事应是重跑短程数值基线与吞吐基线,而不是直接续跑昂贵训练。这样才能知道速度变化来自 kernel,精度变化来自策略,还是数据本身发生了变化。
还应把这些元数据写进 checkpoint 清单,使恢复任务能够拒绝不兼容的精度配置,而不是在数小时后以 NaN 形式暴露。可诊断性本身就是混合精度系统的一项能力。
共同主题是:硬件格式越窄,软件越需要知道张量的数值角色。格式提供可能性,尺度管理、累加方式和回退路径才决定能否稳定训练。
误解一:BF16 比 FP16 精度更高。 BF16 指数范围更大,更少 overflow;但尾数只有 7 位,1 附近分辨率更粗。应说通常更稳定,不是处处更精确。
误解二:loss scaling 能修复所有 NaN。 它主要防 FP16 小梯度下溢。前向激活溢出、除零、softmax 不稳定、学习率过大,都不会被它根治。
误解三:scale 越大越好,而且一定大于 1。 scale 太大会溢出。PyTorch 允许动态 scale 降到 1 以下;应观察 found-inf 和跳步频率。
误解四:用了 autocast 就无需关心 dtype。 自定义算子、显式 dtype、原地运算和策略表外 op 可能保持输入 dtype。关键归约与 loss 仍需审计。
误解五:AMP 会把训练显存直接砍半。 激活可能下降,FP32 优化器状态仍在。账本还包括梯度、通信 buffer、workspace 与碎片。
误解六:FP8 是 BF16 的无痛升级。 FP8 依赖尺度、amax 历史、格式选择和硬件 kernel;错误配置可能不 NaN,却悄悄降低收敛质量。
先运行文末脚本,修改 loss_scaling_sim.py 中 scale 为 1、128、1024、65536。预期 scale 变大时更多微小梯度被保留;继续增大并加入大梯度后会触发 FP16 inf。这正是动态 scale 在 underflow 与 overflow 之间寻找窗口的原因。
在支持 BF16/FP16 的 GPU 上做四组短跑:FP32、BF16 autocast、FP16 autocast 不带 scaler、FP16 autocast 带 scaler。固定随机种子和 batch,记录 200 步吞吐、峰值显存、loss、全局梯度范数、跳步数。预期 FP16 无 scaler 最容易出现零梯度;BF16 与 FP16 加 scaler 更接近 FP32,但结论取决于模型分布。
定位首个异常层时,给模块注册 hook,输出有限值比例、绝对最大值和 dtype。若前向先出现 inf,优先检查指数、归一化、attention score 与输入范围;若反向先异常,再检查 scale、梯度裁剪顺序和自定义 backward。
最后验证累加精度:构造一个大数和大量小数的向量,用 FP32、模拟 BF16 顺序累加、分块 FP32 累加分别求和。预期低精度顺序加法吞掉更多小项;这解释为什么 GEMM 输入低精度不等于 accumulator 也该低精度。
读完这篇可以继续看:
可靠的 AMP 心智模型不是“半精度开关”,而是一张数值预算表:每个张量需要多大范围、多少有效位,在何处累加、何时舍入、出现异常如何回退。把这张表画清楚,NaN 就从玄学变成可以定位的工程问题。
09 节用到的脚本全文如下(float_formats.py、loss_scaling_sim.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。
#!/usr/bin/env python3
"""Only the Python standard library is required."""
import math
import struct
def fp16(x: float) -> float:
try:
return struct.unpack("e", struct.pack("e", x))[0]
except OverflowError:
return math.copysign(math.inf, x)
def bf16(x: float) -> float:
"""Round an IEEE FP32 value to BF16, ties-to-even, then widen for printing."""
bits = struct.unpack(">I", struct.pack(">f", x))[0]
if (bits & 0x7F800000) != 0x7F800000:
bits += 0x7FFF + ((bits >> 16) & 1)
return struct.unpack(">f", struct.pack(">I", bits & 0xFFFF0000))[0]
FORMATS = (
("FP16", 5, 10, 2.0**-14, (2 - 2.0**-10) * 2.0**15, 2.0**-10),
("BF16", 8, 7, 2.0**-126, (2 - 2.0**-7) * 2.0**127, 2.0**-7),
("FP32", 8, 23, 2.0**-126, (2 - 2.0**-23) * 2.0**127, 2.0**-23),
)
print("format exp frac min_normal max_finite epsilon_at_1")
for name, exp, frac, low, high, eps in FORMATS:
print(f"{name:5s} {exp:3d} {frac:4d} {low:.4e} {high:.4e} {eps:.4e}")
values = (1.0001, 0.00001, 100000.0)
print("\nvalue -> FP16 | BF16")
for value in values:
print(f"{value:g} -> {fp16(value):g} | {bf16(value):g}")
#!/usr/bin/env python3
"""Show how FP16 loss scaling rescues tiny gradients; stdlib only."""
import math
import struct
def fp16(x: float) -> float:
try:
return struct.unpack("e", struct.pack("e", x))[0]
except OverflowError:
return math.copysign(math.inf, x)
gradients = [1e-8, 3e-8, 1e-7, 1e-6, 1e-5]
scale = 1024.0
direct = [fp16(g) for g in gradients]
scaled_then_unscaled = [fp16(g * scale) / scale for g in gradients]
print(f"shape=({len(gradients)},), loss_scale={scale:g}")
print("gradient direct_fp16 scaled_fp16/unscaled")
for g, raw, rescued in zip(gradients, direct, scaled_then_unscaled):
print(f"{g:11.3e} {raw:11.3e} {rescued:11.3e}")
print(f"nonzero: direct={sum(x != 0 for x in direct)}/{len(direct)}, scaled={sum(x != 0 for x in scaled_then_unscaled)}/{len(direct)}")
"""Regenerate this article's deterministic teaching figure (numpy + matplotlib)."""
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
OUT=Path(__file__).resolve().parents[1]/"figures"
OUT.mkdir(exist_ok=True)
plt.rcParams.update({"font.sans-serif":["PingFang SC","Arial Unicode MS","DejaVu Sans"],"axes.unicode_minus":False})
from float_formats import fp16,bf16
x=np.geomspace(1e-9,1e5,1500)
y16=np.array([fp16(float(v)) for v in x]);yb=np.array([bf16(float(v)) for v in x])
fig,ax=plt.subplots(figsize=(8,4.8))
for name,y in [("FP16",y16),("BF16",yb)]:
err=np.abs(y-x)/x;err[~np.isfinite(err)]=np.nan
ax.loglog(x,np.maximum(err,1e-12),label=name,alpha=.75)
ax.axvline(65504,color="grey",ls=":",label="FP16 max finite")
ax.set(xlabel="Positive input value",ylabel="Relative rounding error",title="Quantization range and precision (nearest-even)",ylim=(1e-8,2))
ax.legend();ax.grid(alpha=.25);fig.tight_layout()
fig.savefig(OUT/'precision_error.png',dpi=170)
plt.close(fig)
print(OUT/'precision_error.png')
更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」
每日更新 · 论文精选 · 深度解读 · 技术脉络
微信搜索 人工智能炼丹君 或扫描下方二维码关注

评论 (0)