27届大模型面试准备(六十二):大模型数值精度与混合精度工程------BF16/FP8、AMP、Loss Scale 与精度对齐
引言
接(五十)训练稳定性里提到的 BF16 与 loss spike、(十九)量化。本文把"精度"作为工程主线串起来:训练为什么用 BF16、推理为什么上 FP8、混合精度怎么写才不炸、精度对齐怎么验。这是多模态大模型训练 / 推理岗绕不开的基础题,也是面试官判断你"踩过坑"还是"只会调 API"的分水岭。
一、浮点格式对比
| 格式 | 总位 | 符号 | 指数 | 尾数 | 动态范围 | 典型用途 |
|---|---|---|---|---|---|---|
| FP32 | 32 | 1 | 8 | 23 | ~1e±38 | 主权重 / 累加 |
| TF32 | 19* | 1 | 8 | 10 | ~1e±38 | tensor core 矩阵乘 |
| FP16 | 16 | 1 | 5 | 10 | ~1e±4 | 早期训练 / 推理 |
| BF16 | 16 | 1 | 8 | 7 | ~1e±38 | 训练主流 |
| FP8(E4M3) | 8 | 1 | 4 | 3 | ~1e±4 | 推理 / 前向 |
| FP8(E5M2) | 8 | 1 | 5 | 2 | ~1e±15 | 梯度 |
要点:BF16 与 FP32 指数位相同(都是 8 位),所以动态范围一致------梯度再大也不溢出,这是它取代 FP16 成为训练默认格式的根本原因。FP16 尾数多但范围小,容易 overflow / underflow。
二、为什么 BF16 成为训练主流
- 梯度幅值跨多个数量级,FP16 范围 ~6e4 不够,BF16 ~3e38 稳。
- 矩阵乘在 tensor core 上用 TF32/BF16 算力远高于 FP32。
- 代价:尾数只有 7 位,累加误差更大 -> 用 FP32 做"主权重 / master weights"和归约累加来补偿。
一句话:范围决定能不能训,尾数决定训得准不准。BF16 用"范围换尾数",再用 master weights 补回精度。
三、训练混合精度(AMP)
python
import torch
# 方式一:autocast + GradScaler(FP16 场景;BF16 通常不需要 scaler)
scaler = torch.cuda.amp.GradScaler()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
logits = model(x)
loss = crit(logits, y)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
注意:BF16 通常不需要 GradScaler(范围大、不易溢出);FP16 才必须。这是面试高频坑------很多人把 FP16 的 loss scale 套路套到 BF16 上,反而引入不必要的不稳定。
更稳的做法------主权重保持 FP32:
python
# 主权重保持 FP32,每步 cast 到 BF16 计算,梯度回写 FP32
for p in model.parameters():
if p.grad is not None:
p.grad = p.grad.to(torch.bfloat16) # 或保持 FP32 累加后再更新
四、梯度缩放(Loss Scale)
FP16 下梯度很小(如 1e-7)会被舍入成 0 -> 权重不更新。解决:把 loss × 2^k 放大,反向后再 ÷ 2^k 还原。
-
静态:固定 scale(简单但易溢出)。
-
动态:连续几次梯度含 Inf/NaN 就减半,连续正常就翻倍(GradScaler 默认)。
python
# 动态 loss scale 的简化逻辑
scale = 2**15
for step, (x, y) in enumerate(loader):
with autocast(fp16):
loss = model(x)
loss = loss * scale
loss.backward()
if any_nan(model):
scale *= 0.5
zero_grad(); continue
clip_and_step()
scale = min(scale * 2, 2**24)
五、推理量化:FP8 / INT8 / W8A8
-
W8A8:权重与激活都 INT8,per-channel 量化权重、per-tensor 量化激活,用 SmoothQuant 把激活离群值搬到权重。
-
FP8(E4M3):前向用,几乎无损,H100 / 华为昇腾均支持;KV Cache 也可用 FP8 省显存。
-
KV Cache 量化:把历史 KV 从 BF16 压到 INT8/FP8,长上下文场景省 30%~50% 显存。
权重 BF16 --[量化]--> INT8 / FP8 --[算子]--> 输出
^ |
scale / zero v
激活 BF16 --[Smooth]--> INT8 ---------------+--[反量化]
六、精度对齐与溢出检测
- 对齐:新精度实现 vs 参考 FP32 实现,在固定种子的小 batch 上比对 logits / 梯度,允许 1e-2 ~ 1e-3 相对误差。
- 溢出钩子:注册 backward hook,发现 Inf/NaN 立即 dump 张量名与步号,避免模型"静默烂掉"。
python
def nan_hook(module, inp, out):
if isinstance(out, torch.Tensor) and not torch.isfinite(out).all():
raise RuntimeError(f"NaN/Inf at {module.__class__.__name__}")
for m in model.modules():
m.register_forward_hook(nan_hook)
七、面试速答
Q:BF16 和 FP16 训练,谁更易出现 NaN?
A:FP16(范围小,梯度 / 激活易溢出);BF16 范围够大,基本靠 master weights 控误差。
Q:为什么推理上 FP8 比 INT8 更稳?
A:FP8 有指数位,对离群值不敏感,无需复杂 per-channel 校准;INT8 需 SmoothQuant 等处理激活离群值,校准不当就掉点。
八、高频追问清单
- 为什么 tensor core 上 TF32 比 FP32 快这么多?
- master weights 为什么用 FP32 而不是 BF16?
- 动态 loss scale 的"翻倍 / 减半"阈值怎么定?
- per-tensor 与 per-channel 量化差异及适用场景?
- KV Cache 量化对长上下文吞吐的影响?
- 多模态(视觉 encoder 输出)量化要注意什么?
- 如何自动化发现精度回归(precision regression)?
- 混合精度下,哪些算子必须留在 FP32(softmax / 归约 / layernorm)?