先泼一盘冷水:AMP 这个名字在技术圈里经常撞车。搞嵌入式的人,比如最近在调 RK3506,看到 AMP 第一反应是非对称多处理器,满脑子都是核间通信和中断;但站在深度学习训练这一侧,AMP 基本默认指自动混合精度(Automatic Mixed Precision)。这篇只聊自动混合精度,不碰嵌入式场景。文章面向正在被显存和训练速度卡脖子的同学,也适合想搞清楚 autocast 和 GradScaler 到底替你做了什么的读者。我会把底层原理、代码改造、性能收益、常见坑一次性讲透,尽量说人话,不让大家读完还是一头雾水。
1. 混合精度在混合什么:FP16、BF16 和 Tensor Core 的简单模型
1.1 三种数值格式的关键差异
要理解自动混合精度,第一步得先搞清楚计算机里常用的几种浮点格式长什么样。很多人把 FP16 理解成“把 FP32 截短一点”,这个说法方向没错,但实际差别比想象中微妙。
FP32 也叫单精度浮点数,用 1 个符号位、8 个指数位、23 个尾数位表示一个数。FP16 是半精度,用 1 个符号位、5 个指数位、10 个尾数位。BF16 是脑浮点,用 1 个符号位、8 个指数位、7 个尾数位。
| 格式 | 符号位 | 指数位 | 尾数位 | 最大有限值 | 最小正规数 | 相对精度 |
|---|---|---|---|---|---|---|
| FP32 | 1 | 8 | 23 | 约 3.4e38 | 约 1.18e-38 | 高 |
| FP16 | 1 | 5 | 10 | 65504 | 约 6.10e-5 | 中 |
| BF16 | 1 | 8 | 7 | 约 3.4e38 | 约 1.18e-38 | 低 |
这里的核心区别是动态范围和精度。FP16 因为指数位只有 5 位,最大只能表示到 65504,超过就会变成 Inf;最小正规数是 6.1e-5 左右,很多小的梯度值天然比这个还小,直接存成 FP16 会直接变成 0,这就是“下溢”。BF16 保留了和 FP32 一样的 8 位指数,所以动态范围很安全,但尾数只有 7 位,精度比 FP16 还要低,很多训练细节直接被抹平。
你可以这样理解:FP32 是一张能记大数小数的账本,FP16 是只能在 0 到 65504 之间记数的账本,BF16 是账本足够大但只允许你写 6 位有效数字。AMP 要做的就是在不同场景里选择用哪本账本,既要把速度提上去,又不能把账算崩。
1.2 为什么不能直接全转成 FP16
有个很常见的误解:自动混合精度就是把模型所有权重、梯度、激活值全转成 FP16。如果真这么干,模型大概率在训练早期就开始发散,原因主要有三个。
第一,某些算子对数值范围极其敏感。比如 Softmax 要做指数运算,输入里一旦出现大数值,经过 FP16 的有限范围后很容易溢出;LayerNorm、BatchNorm 这类归一化算子内部也有求均值、方差的过程,如果强制用 FP16 计算,统计精度会明显下降,训练曲线经常直接起飞。第二,反向传播里的梯度很多是小数值,直接落在 FP16 的表示范围之外,下溢成 0 之后,浅层参数几乎学不动。第三,虽然现代 GPU 上有 Tensor Core 可以用 FP16 做矩阵乘,但并不是所有算子都有对应的 FP16 实现,如果某些自定义算子只支持 FP32,全转 FP16 之后连跑都跑不起来。
所以混合精度的正确姿势是:把支持低精度运算、对精度不敏感、计算量又大的部分,比如卷积、矩阵乘法、线性层,交给 FP16 去算;把对数值范围敏感、容易崩的部分,比如归一化、指数类操作,继续保留 FP32;梯度方面再用专门的机制防止下溢。AMP 的价值恰恰在于,这套调度由框架自动完成,不需要开发者手动去给每个算子做分类。
1.3 Tensor Core 到底快在哪
NVIDIA 从 Volta 架构开始引入 Tensor Core,Tensor Core 擅长执行 FP16 输入、FP32 累加的矩阵乘。简单说,它允许你用两个半精度矩阵做乘法,内部累加时用 FP32,最后输出的精度损失可控。因为单次硬件吞吐大幅提升,矩阵乘类算子在半精度下往往能跑到 FP32 的两倍甚至更高,这也是 AMP 在 Transformer、卷积网络这类矩阵乘密集型模型上收益最明显的原因。
不过要注意,Tensor Core 带来的提速不是免费的。它要求数据在参与特定运算时走 FP16 路径,如果模型里大量算子没有落到 Tensor Core,那么 AMP 的收益就会被稀释。这也是为什么后面我反复强调,用 AMP 之前先确认你的模型计算主体是不是卷积、矩阵乘和注意力,这部分占比越大,收益越大。
2. AMP 核心机制拆解:autocast 与 GradScaler 的真正分工
2.1 autocast:自动给你挑选计算精度
PyTorch 里 AMP 体系主要由两个组件构成,一个是torch.autocast,一个是torch.amp.GradScaler。这两个各管一摊,缺一不可。
autocast是一个上下文管理器。进入这个上下文之后,PyTorch 会拦截参与自动类型转换的算子,按照一张预定义的分派表来决定当前运算用什么精度执行。并不是把所有算子都压成 FP16,而是分成了三类。第一类是明确可以低精度执行的算子,比如卷积、线性层、矩阵乘,会优先使用 FP16。第二类是必须在 FP32 下执行的算子,比如 Softmax、LayerNorm、BatchNorm 这类归一化相关算子,以及部分数值稳定要求高的算子,自动保持在 FP32。第三类是一些按照输入 dtype 来决定的算子,输入如果是 FP16 就按 FP16 算,输入否则就按 FP32 算。
使用autocast有一个关键认知:它不会改变模型参数本身的 dtype。你在模型定义里如果用nn.Linear,参数默认还是 FP32,只是在 forward 进入 autocast 区域之后,参与计算的输入会被临时转成 FP16,计算完成后输出类型按规则恢复。所以不要指望跑完一个 AMP 前向传播后,模型权重真的变成 FP16 存起来了。
with torch.autocast(device_type="cuda", dtype=torch.float16): output = model(inputs) loss = criterion(output, targets)前面提到过,BatchNorm 在 FP16 下不稳定,但在实际使用中,如果模型里有 BatchNorm,往往还有个更麻烦的点:BatchNorm 在训练和推理模式下统计均值方差的行为不同,再加上 AMP 的类型转换,经常会出现训练时正常、推理时结果对不上的情况。我的建议是,BatchNorm 特别多的模型先用小批量数据跑通后再上 AMP,不要一上来就大改代码。
2.2 GradScaler:给梯度加动态保险丝
autocast只解决前向传播和反向传播过程中算子精度分配的问题,它不管梯度下溢。FP16 的最小正规数是 6e-5 左右,很多梯度的绝对值比这个还小,如果不做处理,梯度会在反向传播中变成 0,参数不再更新。GradScaler 就是干这个的。
GradScaler 的核心思路是:在 loss 反向传播之前,先把 loss 乘上一个缩放因子,再调用backward()。因为梯度是 loss 的导数,loss 被放大之后,梯度也跟着被放大了,这样原本小于 FP16 表示范围的小梯度就能落到可表示区间里。等到反向传播完成、优化器更新之前,再把梯度除以缩放因子还原回去。
scaler = torch.amp.GradScaler("cuda") ... scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()注意,scaler.step(optimizer)内部会先检查这一个迭代的梯度里有没有 Inf 或者 NaN,如果有,就跳过这一轮参数更新,然后scaler.update()会把缩放因子按照动态策略调小;如果没有,就正常更新参数,并且每隔一定步数尝试把缩放因子调大一点。默认初始缩放因子通常是 65536,对应 FP16 的安全边界,乘以 2 倍增长,除以 2 倍回退,这样的设计保证了缩放因子能随着训练动态调整到合适区间。
2.3 为什么不直接把所有梯度手动放大
有人会问,既然小梯度会下溢,那我手动把梯度全部乘以一个大数再反传,不就行了?这一步看起来没问题,但实际很容易出偏差,因为手动放大梯度会导致 loss 变化范围变得很大,如果你后面代入了别的精度转换、混合了多个优化器,或者把梯度裁剪加进来,很容易把顺序搞乱。GradScaler 并不是简单放大 loss,它把自己的状态和 optimizer、backward 流程耦合在一起,保证全部逻辑都按约定顺序跑。
你可以把 GradScaler 想象成一个会自动调节的放大镜:看清小字的时候把放大倍数调大,发现画面过曝就调小,稳定一段时间后又尝试放大一点。这个动态调整过程是被训练迭代状态驱动的,省心且不容易翻车。
还有一个值得注意的点:如果你的模型用的是 BF16,GradScaler其实可以不开。因为 BF16 的动态范围和 FP32 一样,小梯度不会下溢到 0,不存在 FP16 那种溢出和生产风险。很多新架构、新显卡上,大家更倾向于用 BF16 而不是 FP16,就是为了省去 GradScaler 这一层复杂度。
3. 实操:把一个普通 PyTorch 训练循环改造为 AMP 版
3.1 最小改动模板:从普通循环到 AMP 循环
假设你有一个非常标准的 PyTorch 训练循环,原来的代码大概是这样的:
for batch in dataloader: inputs = batch["inputs"].cuda() labels = batch["labels"].cuda() optimizer.zero_grad() outputs = model(inputs) loss = criterion(outputs, labels) loss.backward() optimizer.step()改成 AMP 版本,只需要加三件东西。第一,在训练循环外初始化一个 GradScaler;第二,把 forward 包进torch.autocast;第三,把loss.backward()改成scaler.scale(loss).backward(),把optimizer.step()改成scaler.step(optimizer),并在每轮之后调用scaler.update()。
scaler = torch.amp.GradScaler("cuda") for batch in dataloader: inputs = batch["inputs"].cuda() labels = batch["labels"].cuda() optimizer.zero_grad() with torch.autocast(device_type="cuda", dtype=torch.float16): outputs = model(inputs) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()改动量就这么大。这个模板在多数 CNN、Transformer 模型上可以直接用,不用手工改任何一个网络的 forward。唯一要额外留意的是 loss 的记录和传递,因为 autocast 区域内部得到的 loss 可能是低精度类型,如果你习惯用loss.item()记录训练日志,最好先float(loss.detach().float())转成 Python float,否则日志里的数值可能因为精度被截断,看起来不太正常,但不影响训练本身。
3.2 加入梯度裁剪的正确写法
AMP 训练里最容易搞错的顺序问题就是梯度裁剪。普通训练里你可能会这么写:
loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step()但在 AMP 下不能这么写。因为这时候梯度已经被 GradScaler 放大了,如果你直接对放大后的梯度做裁剪,裁剪幅度是按放大后的尺度算的,等优化器更新时又会被缩回去,前后尺度不对齐,clip 几乎等于没生效。
正确的做法是先调用scaler.unscale_(optimizer)把梯度还原,再裁剪,最后再scaler.step(optimizer):
scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) scaler.step(optimizer) scaler.update()unscale_这个名字很直白,就是把缩放因子去掉。调用之后,梯度恢复成原始尺度,之后做裁剪、查看梯度范数都是可解释的。如果你忘了这一步,最常见的问题是:loss 看起来在正常下降,但模型学得很慢,甚至几乎不收敛。
3.3 多优化器场景:GAN、多任务训练怎么写
在 GAN 或多任务训练里,经常会出现两个优化器,比如一个优化生成器,一个优化判别器。此时 GradScaler 的写法有一些小讲究。最稳妥的做法是,backward 都完成之后,先把所有优化器的梯度都 unscale 掉,然后逐个调用scaler.step(optimizer),最后统一调用一次scaler.update()。
# 这里 d_loss 和 g_loss 分别是判别器和生成器的 loss scaler.scale(d_loss).backward() scaler.scale(g_loss).backward() scaler.unscale_(optimizer_d) scaler.unscale_(optimizer_g) scaler.step(optimizer_d) scaler.step(optimizer_g) scaler.update()为什么要先把两个优化器都 unscale?因为scaler.step内部会检查梯度里有没有 Inf/NaN,如果只 unscale 其中一个,另一个优化器的梯度仍然是缩放后的尺度,检查动作就不完整。多个优化器叠加时,宁可多写两行 unscale,也不要省。
3.4 保存断点和恢复训练时,scaler 状态一定要带上
很多人保存模型只保存model.state_dict()和optimizer.state_dict(),但用 AMP 训练时,GradScaler 也有自己的状态,包括当前的缩放因子、已经连续多少个 iteration 没有出现 Inf/NaN。如果断点恢复时漏掉 scaler 的 state,缩放因子会重置到初始值,训练虽然不会立刻崩,但动态调整的节奏被打断,后面可能反复触发 NaN 检查,白白浪费时间。
# 保存 torch.save({ "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scaler": scaler.state_dict(), }, "checkpoint.pt") # 恢复 checkpoint = torch.load("checkpoint.pt") model.load_state_dict(checkpoint["model"]) optimizer.load_state_dict(checkpoint["optimizer"]) scaler.load_state_dict(checkpoint["scaler"])4. 踩坑实录:7 个训练中容易翻车的 AMP 问题
4.1 Loss 突然变成 NaN,不一定是 AMP 的锅
AMP 训练最常见的问题是 loss 跑着跑着变成 NaN。遇到这种情况,我第一件事不是关掉 AMP,而是先去判断问题出在哪一层。可以把scaler.get_scale()打出来看,如果缩放因子已经被连续调小到很小,说明之前有梯度溢出被 GradScaler 抓到了;如果缩放因子始终保持不变,loss 却突然 NaN,基本可以排除 GradScaler 的锅,主要矛盾在模型本身的学习率、初始化和数据。
有个非常实用的排查技巧:把dtype=torch.float16换成dtype=torch.bfloat16再跑一次。如果 BF16 下 loss 正常,那问题多半是 FP16 动态范围太小导致某个中间值溢出了;如果 BF16 下也 NaN,那就要回头检查模型结构、学习率 scheduler、数据 pipeline,别在一个无关的位置死磕 AMP。
4.2 模型不更新或更新幅度异常,先看 unscale 顺序
这个问题我在 3.2 里提过,但实际群里问到的人实在太多,值得单拎出来再强调一次。训练过程中如果发现 loss 下降得非常慢,或者某个模块的梯度范数一直异常大,先用下面这个组合排查:确认loss.backward()用的是scaler.scale(loss).backward();确认梯度裁剪之前调用了scaler.unscale_(optimizer);确认optimizer.zero_grad()没有被放在scaler.scale(loss).backward()之后漏掉。很多人改了 AMP 后仍然在用老的loss.backward(),结果梯度没有被缩放,小梯度全部下溢,模型看起来“没死”,但效果怎么都上不去。
4.3 autocast 和自定义算子冲突
如果你的模型里有自己写的 CUDA 扩展、自定义 autograd Function,AMP 不一定能自动处理。PyTorch 的内置算子大多在 autocast 分派表里有覆盖,但自定义算子通常需要手动实现 autocast 的 cast 函数,否则它可能会按输入 dtype 直接执行,得到一个意想不到的 FP16 结果。
如果你用的自定义算子本身只支持 FP32,进 autocast 区域后会把 FP16 输入转回 FP32,甚至报“not implemented for Half”的错误。这种情况,一个比较省心的做法是:在自定义算子的backward里显式 cast 到需要的精度,或者在 forward 前后手动to(torch.float32)隔离,避免让自定义算子掺和进 fp16 自动降精度的流程里。
4.4 带有 BatchNorm 的模型在推理阶段结果异常
前面提到了 BatchNorm 在 AMP 训练下表现还可,但部署推理时经常会有坑。训练时 BatchNorm 会在每个 batch 上计算统计量并更新 running mean 和 running var,推理时直接使用 running 统计量。由于训练时 forward 内部发生了一轮 FP16 临时转换,BatchNorm 的 running 统计量是在 FP32 上维护的,这没问题,但如果你推理时把整个模型.half()或者漏掉 autocast,前后两端的精度状态不一致,输出分布就可能偏移。
我的习惯是:训练阶段用torch.autocast,推理阶段也用同样配置的torch.autocast,不要一边用 AMP 训练、一边用纯 FP16 模型推理,两边精度状态必须对齐。如果模型里 BatchNorm 特别重,又希望推理彻底省心,更推荐用 ONNX 导出后做量化,而不是纯靠 AMP 硬扛。
4.5 分布式训练里 DDP 和 GradScaler 的协作问题
用 DistributedDataParallel 跑 AMP 时,多数情况没有额外问题,因为梯度同步发生在反向传播过程中,GradScaler 是在反向传播完成后才介入 unscale 的,所以不会破坏梯度同步。但需要注意,GradScaler 检测 Inf/NaN 时是基于当前 rank 的梯度判断的,如果某个 rank 出现 NaN,这一个迭代会跳过 optimizer.step,其他 rank 的步数会不一致。虽然 DDP 本身不要求每个 rank 的迭代次数完全一致,但为了日志和评测对齐,最好统一判断条件,或者在主进程里收集scaler.get_scale()做监控。
多卡训练里另一个容易踩的坑是torch.nn.SyncBatchNorm这类同步逻辑,它需要额外的通信开销,如果和 AMP 混用,建议先单独跑通小规模验证,确认通信量和稳定性没问题再放大。
4.6 梯度累积场景下 scaler.update 的节奏
梯度累积时,一个容易混乱的点是:是不是每个 mini-batch 都要调用scaler.update()?我的经验是,在累积模式下,每个 mini-batch 仍然要调用scaler.scale(loss).backward(),但在真正执行 optimizer.step 的那个一个迭代里调用scaler.step(optimizer)和scaler.update()。如果你在每一个累积 mini-batch 都调用scaler.update(),缩放因子的调整频率会被放大,GradScaler 的步数计数会乱掉,动态缩放策略就不准确了。
为了省事,有人会干脆每次 backward 都调用 scaler.step 一次,再配一个空优化器,这有点自欺欺人。更清晰的做法是把梯度累积主循环拆成“累积阶段”和“更新阶段”,只在更新阶段调用scaler.step和scaler.update。
4.7 旧版 PyTorch 的 API 差异
不同 PyTorch 版本的 AMP 接口有细微差别。老版本里很多人用from torch.cuda.amp import autocast, GradScaler,新版本更推荐torch.amp.GradScaler("cuda")和torch.autocast(device_type="cuda")。如果你的项目还打在旧版接口上,可以先检查一下用的 PyTorch 版本,把接口统一升级到新版写法,这样后面再迁移到 DeepSpeed、Megatron 这类框架时会顺滑一些。
5. 性能收益怎么看:显存、吞吐和收益边界
5.1 显存收益来自哪一层
很多文章把 AMP 的显存收益简单说成“权重减半”,这其实不够准确。原生 PyTorch AMP 运行时,模型参数仍然保持 FP32,不像 Apex 的 O2 模式会主动把权重转半。所以如果你模型里最占显存的是权重本身和 Adam 优化器状态,AMP 给你省的主要是激活值和一部分临时张量,静态权重占用并不会减半。
激活值减半带来的收益在长序列、大 batch 的模型上非常明显。比如 Transformer 训练时,中间激活值经常比权重还占显存,FP16 激活值会比 FP32 省下一大块。所以在实际项目中,AMP 带来的显存收益通常表现为能让你把 batch size 调大到原来的 1.2 到 1.8 倍,而不是直接把模型文件体积减半。
5.2 速度收益怎么看才科学
提升速度是 AMP 的核心卖点,但很多人测试方式不对,跑两三个 step 就下结论。正确做法是让模型先跑几十个迭代做 warmup,把 CUDA kernel 预热、显存分配、缓存状态全激活后再计时;统计时间时用稳定段落的平均值,而不是第一轮时间。还可以用torch.cuda.max_memory_allocated()记录显存峰值,用torch.cuda.synchronize()保证计时准确,避免异步计算把上一轮的 kernel 时间算到下一轮。
一般来说,矩阵乘占比高、Tensor Core 支持好的模型收益最理想,实测里很多 Transformer 类任务能提升 30% 到 80% 的训练吞吐。反过来,如果你的模型全是小算子、控制流、自定义逻辑,或者你的显卡本身不具备 Tensor Core,AMP 的提速效果会非常有限,有时候更慢。
5.3 什么时候不值得上 AMP
不上 AMP 的情况也很明确。第一,模型特别小,显存不紧张,训练时间里 CPU 数据加载和预处理占了主导,这时上 AMP 属于给水桶换一个更大的出水口,但进水口没变,收益可以忽略。第二,模型里有大量必须先保持 FP32 的自定义操作,AMP 能覆盖的计算比例太低,提速被 FP32 路径拖慢。第三,你的场景对 bit 级复现有要求,FP32 和 AMP 的结果不会完全一致,因为低精度运算本身带来取舍,哪怕只是多跑几个迭代,loss 曲线也会有细微差别。
6. 再往前一步:BF16、推理优化和断点恢复
6.1 BF16 正在成为很多新场景的首选
FP16 在训练时的动态范围问题让不少人头疼,而 BF16 因为指数位和 FP32 一样,基本不存在溢出困难,所以在新一代硬件上,很多训练流程直接把 FP16 换成了 BF16。使用方式也很简单,把dtype从torch.float16换成torch.bfloat16,并且不需要 GradScaler:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16): outputs = model(inputs) loss = criterion(outputs, labels)BF16 的缺点是尾数位少,精度低,如果你测试发现验证精度或者 loss 在 BF16 下掉了太多,可以回到 FP16 配合 GradScaler 的方案。另一种常见策略是在训练前期用 FP32,等模型进入稳定区间后再切换到 BF16,这种按阶段调精度的思路在不少项目里都能兼顾收敛速度和稳定性。
6.2 推理阶段用 AMP 的正确姿势
有些人在推理时不用autocast,而是直接把模型.half(),然后发现效果变差甚至报错。根因在于并不是所有算子都适合半精度,直接 all half 相当于把模型所有部分都压成 FP16,缺少了 autocast 的算子级调度。如果你只想用最少的代码提升推理速度,建议用和训练一模一样的 autocast 包住推理块:
model.eval() with torch.inference_mode(): with torch.autocast(device_type="cuda", dtype=torch.float16): outputs = model(inputs)这个方案比手动.half()稳得多,而且具备自动回退 FP32 的能力。目前很多部署框架如 TensorRT 原生支持动态精度分配,如果要做极致部署,不一定要继续用 PyTorch 的 AMP 逻辑。
6.3 我踩过几次坑之后的个人体会
把整个 AMP 流程理顺之后,你会发现它其实并不神秘,本质上就是两件事:autocast 管算子在什么精度下执行,GradScaler 管梯度不下溢。最难的部分永远不是 API 怎么调用,而是当你面临一个具体模型时,能不能准确判断出到底是哪个环节出了精度问题。我现在的习惯是,任何新模型接入 AMP 前都先跑一个 50 步的小验证,对比 FP32、FP16、BF16 三条曲线的 loss 变化趋势;确认稳定之后再做完整训练。这条习惯帮我省下来很多试错时间,也让后面调学习率、改结构时有了一个可以横向对比的基准。AMP 不是万能特效药,但它真的能在模型足够大、显存足够紧张的时候,把你从“显存不够用”的焦虑里救出来,值得你花一个下午把它彻底搞明白。