ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch自动混合精度AMP原理与实战:显存减半、训练提速

PyTorch自动混合精度AMP原理与实战:显存减半、训练提速 先泼一盘冷水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 个尾数位。格式符号位指数位尾数位最大有限值最小正规数相对精度FP321823约 3.4e38约 1.18e-38高FP16151065504约 6.10e-5中BF16187约 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 CoreTensor 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_typecuda, dtypetorch.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 想象成一个会自动调节的放大镜看清小字的时候把放大倍数调大发现画面过曝就调小稳定一段时间后又尝试放大一点。这个动态调整过程是被训练迭代状态驱动的省心且不容易翻车。还有一个值得注意的点如果你的模型用的是 BF16GradScaler其实可以不开。因为 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_typecuda, dtypetorch.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_norm1.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_norm1.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 的锅主要矛盾在模型本身的学习率、初始化和数据。有个非常实用的排查技巧把dtypetorch.float16换成dtypetorch.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 FunctionAMP 不一定能自动处理。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_typecuda)。如果你的项目还打在旧版接口上可以先检查一下用的 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 CoreAMP 的提速效果会非常有限有时候更慢。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并且不需要 GradScalerwith torch.autocast(device_typecuda, dtypetorch.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_typecuda, dtypetorch.float16): outputs model(inputs)这个方案比手动.half()稳得多而且具备自动回退 FP32 的能力。目前很多部署框架如 TensorRT 原生支持动态精度分配如果要做极致部署不一定要继续用 PyTorch 的 AMP 逻辑。6.3 我踩过几次坑之后的个人体会把整个 AMP 流程理顺之后你会发现它其实并不神秘本质上就是两件事autocast 管算子在什么精度下执行GradScaler 管梯度不下溢。最难的部分永远不是 API 怎么调用而是当你面临一个具体模型时能不能准确判断出到底是哪个环节出了精度问题。我现在的习惯是任何新模型接入 AMP 前都先跑一个 50 步的小验证对比 FP32、FP16、BF16 三条曲线的 loss 变化趋势确认稳定之后再做完整训练。这条习惯帮我省下来很多试错时间也让后面调学习率、改结构时有了一个可以横向对比的基准。AMP 不是万能特效药但它真的能在模型足够大、显存足够紧张的时候把你从“显存不够用”的焦虑里救出来值得你花一个下午把它彻底搞明白。
返回列表