混合精度训练是什么?为什么能加速深度学习?
简化版
混合精度用 FP16/BF16 执行适合的矩阵计算,同时把关键累加、主权重或敏感算子保留 FP32,以降低显存并利用 Tensor Core。FP16 常需动态 loss scaling 防梯度下溢,BF16 指数范围接近 FP32,通常不需要 scaling。
详细版
-
autocast 根据算子白名单选择 dtype,不是把全模型强制 half。
-
FP16 指数范围窄,小梯度易下溢,loss scale 在反向前放大损失。
-
optimizer step 前要 unscale,再做梯度裁剪和非有限值检查。
-
BF16 尾数更短但范围大,稳定性通常优于 FP16。
-
加速取决于硬件、shape 对齐和是否受访存/非矩阵算子限制。
完整版教学
一、混合精度在范围、精度与吞吐间分工
大矩阵乘对少量舍入误差通常鲁棒,低精度能减少读写并提高专用单元吞吐;归一化、归约和优化器状态对累计误差更敏感,常保留 FP32。
Loss scaling 只改变反向数值表示:放大后的梯度计算完再除回去,理论更新不变;发生 inf/NaN 时跳过该步并降低 scale。
二、底层机制与公式
scaled_loss = loss * S
scaled_grad = grad * S
unscaled_grad = scaled_grad / S
FP16 max ≈ 65504; BF16 exponent range ≈ FP32
三、带数字的推演
真实梯度 1e-8 在 FP16 可能下溢为 0;取 S=65536 后先表示为约 6.55e-4,unscale 后恢复更新尺度。
若放大后溢出,scaler 会跳过更新并减小 S。
四、方案对比
| 方案/对象 | 核心特点 | 代价或边界 |
|---|---|---|
| FP32 | 精度和工具兼容最好 | 显存/吞吐成本高 |
| FP16+scaler | 硬件加速成熟 | 范围窄,需溢出管理 |
| BF16 | 范围大、通常免 scaling | 尾数短且需硬件支持 |
五、执行流程
autocast 前向 -> 计算 FP32 loss -> scaler.scale(loss).backward
-> unscale optimizer -> clip/check grads -> scaler.step -> scaler.update
六、边界条件与工程代价
将模型直接 .half() 会让 LayerNorm、softmax 等敏感运算也落入 FP16,容易出错;优先使用框架 autocast。
低精度节省激活显存,但 Adam 的 FP32 状态可能仍占大头。
要估算总显存,必须把参数、梯度、优化器状态和激活分开。
记忆钩子:把混合精度记成“矩阵低精度跑,归约关键处高精度守,FP16 用 scaler 抬高小梯度”。
七、常见误区与追问
-
误区:混合精度就是所有张量改为 FP16。 敏感算子和累加通常仍用 FP32。
-
追问:BF16 为什么通常不需 loss scaling? 其指数位与 FP32 相同,动态范围更大。
-
误区:裁剪 scaled gradient 也一样。 阈值会作用在放大后的数值,必须先 unscale。
-
追问:为何有时不加速? 算子太小、shape 未对齐或任务受访存和 CPU 限制。
-
追问:如何监控稳定性? 记录 scale、跳步次数、非有限梯度与 FP32 基线指标。
八、加强记忆
把混合精度记成“矩阵低精度跑,归约关键处高精度守,FP16 用 scaler 抬高小梯度”。
顺序必须是 backward、unscale、clip、step、update。