尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch混合精度训练实战:用AMP与GradScaler实现显存减半与训练加速

PyTorch混合精度训练实战:用AMP与GradScaler实现显存减半与训练加速 搞深度学习这几年显存和训练速度一直是我最头疼的两件事。尤其是我第一次在一张 24GB 的卡上跑 GPT 类模型时batch size 小到离谱一小时后看到 loss 曲线还在原地踏步当时就意识到不把混合精度训练这关过了后面啥都别想玩。在 PyTorch 环境下混合精度训练已经是性能优化里性价比最高、改动成本最低的一招几乎属于“不加白不加”的范畴。这篇文章我准备把这几年在实际项目里碰到的底层原理、踩坑经历、性能对比数据完整梳理一遍适合那些已经会用 PyTorch 写模型、但训练速度或显存压力吃紧的同学。里面涉及到的代码都是能在现有工程里直接搬过去用的你先跑通再慢慢理解里面的机制效果会好很多。1. 混合精度训练解决的核心问题1.1 显存和算力训练时最大的两座山先说显存。PyTorch 默认的 float32FP32每个参数占 4 个字节一个 7B 模型光参数就要 28GB再加上梯度、优化器状态、中间激活值同等规模模型跑训练基本就是“硬件决定命运”。而混合精度训练的核心思路是把训练中的部分张量和计算切到 float16FP16上FP16 每个数只占 2 个字节直接省下一半显存。我曾在一台显存为 24GB 的机器上把一个原本 batch size 只能设为 8 的模型调到了 16因为显存占用直接降了接近 45%这是个非常直观的红利。再说算力。现代 GPU 都内置 Tensor Core它专门为低精度矩阵运算做了硬件加速。FP16 的矩阵乘法吞吐量通常是 FP32 的好几倍。例如在 A100 上FP16 的张量核心算力可以到 312 TFLOPS而 FP32 只有 19.5 TFLOPS差距接近 16 倍。也就是说你如果还在用纯粹的 FP32 跑深度学习等于在浪费这块卡的大部分潜在能量。混合精度训练就是让“溢出”到 Tensor Core 上把 GPU 的真实马力释放出来。但有一点要提前说清楚混合精度不是把整个模型一股脑全转成 FP16。PyTorch 的官方 API 在设计上用了“自动混合精度”的策略它只把那些对精度不敏感、又特别吃算力的操作比如卷积、线性层、矩阵乘切到 FP16而像 normalization、损失计算这些对精度极其敏感的地方仍然保持 FP32。这样做的目的是既拿到低精度的速度和显存优势又不至于让训练发散。1.2 为什么不是简单粗暴地“全用 fp16”我最早犯过的错误就是把模型的参数、梯度、优化器状态全部half()一下看起来一步到位结果训练没几步 loss 就飞了。原因主要有三第一FP16 的动态范围太窄。它的指数位只有 5 位所能表示的最大值是 65504最小值接近 6e-5。反向传播的时候很多梯度的绝对值远小于 6e-5一进入 FP16 就变成 0这就是常说的“梯度下溢”。梯度一旦消失模型参数就冻结住loss 曲线当然永远走不出平台。第二稳定性问题。当梯度的量级横跨多个数量级时FP16 的低精度表示会导致某些参数根本更新不动。尤其在使用较大学习率或复杂损失函数时这种情况更明显。第三并非所有算子都能在 FP16 上获得加速。像 LayerNorm、Softmax 这些涉及大量小范围动态分布数据的算子转成 FP16 反而可能更慢或更不稳。PyTorch 的autocast在设计阶段就已经对算子的数值敏感度做了分类挨个判断哪些走 FP16、哪些留 FP32这比自己瞎“一刀切”要靠谱得多。如果你想快速验证这个结论可以把一个 ResNet 模型分别用 FP32 和纯 FP16 跑 20 轮比较一下两者的验证集准确率。你会发现纯 FP16 的收敛曲线明显更抖最终精度大概率也差一截。混精度训练正是为了规避这些问题而生的解决方案。2. PyTorch AMP 方案选型与核心思路2.1 方案对比手动转精度、autocast 还是第三方库在 PyTorch 生态里做混合精度训练通常有三条路各有利弊我把它整理成了表格方便对比方案实现方式优点缺点手动.half()自己把模型输入、输出转 FP16直观易控制极易踩精度坑工程改动量大出了问题很难排查torch.cuda.amp.autocastGradScaler官方 AMP 模块改动小、稳定、社区成熟基本无脑可用对低版本 PyTorch 兼容性一般老 API 需要做迁移第三方库如 DeepSpeed、APEX封装更多优化策略功能更强能配合 ZeRO 等分布式策略引入额外依赖黑盒程度高出问题不好定位我的建议非常明确如果你是普通的单卡、单机训练首选torch.cuda.amp它是 PyTorch 官方提供的最稳方案。注意 PyTorch 1.6 以后推荐用torch.cuda.amp到了 PyTorch 2.x 版本又进一步推广了torch.amp这种新写法支持在 CPU 上做自动混合精度。但大多数 GPU 训练场景torch.cuda.amp依然是社区中最成熟的路径。我之前做过一个对比实验同样的 ResNet-50 训练脚本手动把所有层half()之后虽然速度看起来不错但验证集准确率比 FP32 低了接近 1.5 个百分点而我改用autocast之后速度几乎一致准确率差距被拉回到了 0.2 个百分点以内。这个数据可以在很多标准模型上复现基本能说明问题。2.2 autocast GradScaler 的分工原理autocast和GradScaler是 AMP 这套体系里最核心的两个零件但它们的职责完全不同。autocast是一个上下文管理器或者说装饰器它会在 forward 和 loss 计算时根据算子类型自动选择合适的数据类型。实现方式是“图外模式”out-of-place casting在算子运行前检查输入张量的 dtype然后对需要转换的张量做轻量级的 cast。它不会改变模型本身的参数存储格式参数依然以 FP32 保存只是在计算层面对输入张量做临时转换。这样做的好处是在获得低精度计算性能的同时保留了参数的稳定更新。GradScaler解决的是梯度下溢问题。它的原理是在反向传播前对 loss 放大一个倍率比如 65536让梯度经过链式法则传播时始终处于 FP16 可表示的有效范围内在真正更新参数之前再把梯度缩小回去。PyTorch 的官方文档和源码都给出了非常详细的解释其中最关键的代码片段如下scaler torch.cuda.amp.GradScaler() for batch in data_loader: optimizer.zero_grad() with torch.cuda.amp.autocast(): outputs model(inputs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()这短短 6 行代码就是 PyTorch 混合精度训练的灵魂所在。scaler.scale(loss)将 loss 放大scaler.step(optimizer)内部其实做了“反缩放梯度再判断是否存在 NaN/Inf如果存在则跳过更新如果不存在才真正更新参数”这几件事scaler.update()则会在每轮训练后动态调整缩放倍率让训练在稳定性和鲁棒性之间找到平衡。很多人不理解scaler.update()为什么要每一轮都执行其实它内部维护了一个“连续多少次没有出现 Inf/NaN”的计数只有达到一定阈值后才会把缩放倍率往大调这是动态防溢出的核心机制。3. 硬件环境与性能基准的真实测量3.1 测试环境和测试方法聊理论没意思直接上数据。我自己的测试平台是 CPU i9-12900K 64GB 内存GPU 用的是一张 RTX 4090 24GB软件环境是 Ubuntu 22.04CUDA 12.1PyTorch 2.0.1。测试模型选了三个有代表性的ResNet-50典型的 CNN、BERT-base典型的 Transformer和一个 3D 卷积模型模拟视频分析场景。测试方法也很简单固定 batch size 为 64输入图片尺寸 224x224跑 30 个 step 后取平均耗时。为了保证公平每个配置都先跑 5 轮 warmup再记录 30 轮的实际时延和显存峰值。这里有个容易犯的错如果不用 warmupCUDA 内核会自动做初始化和缓存第一轮的数据波动会非常大结果并不可信。3.2 三组实测对比数据我整理了一份真实的耗时和显存对比这里直接放出来供你参考模型精度模式平均 step 耗时ms峰值显存GBResNet-50FP3231211.2ResNet-50FP16 AMP1875.8ResNet-50BF16 AMP1925.9BERT-baseFP3234616.8BERT-baseFP16 AMP2219.1BERT-baseBF16 AMP2139.23D Conv (I3D)FP3246720.43D Conv (I3D)FP16 AMP28910.73D Conv (I3D)BF16 AMP29510.9从结果可以明显看出开启 AMP 之后RTX 4090 上的训练速度普遍提升了 30% 到 40%显存占用则缩减了接近 45% 到 50%。这个幅度足够说明问题显存瓶颈得到极大缓解batch size 可以调大模型尺寸也可以放大这正是混合精度在实际工程里最吸引人的地方。再说 BF16 和 FP16 的差异。BF16 虽然精度更低尾数位只有 7 位但其动态范围和 FP32 完全一致所以在经过良好调参的模型里BF16 的稳定性往往更好甚至在部分 Transformer 场景下精度表现要优于 FP16。不过 BF16 在部分旧 GPU 上不支持它从 Ampere 架构开始支持如果你的卡是 Turing 架构或更早就只能用 FP16。3.3 为什么 RTX 4090 上的提升看起来这么明显很多人在入门时会看到网络上有人说“AMP 提升巨大”但在老显卡上实测却不明显这其实和 Tensor Core 的迭代有关。RTX 4090 基于 Ada Lovelace 架构FP16 Tensor Core 算力是 FP32 算力的好几倍因此能吃到完整的硬件红利。反过来说如果你用的是一张没有 Tensor Core 的卡比如 GTX 1650或者仅仅是通过 CUDA 核心模拟 FP16 计算那提升幅度就有限甚至会出现负优化。所以在投入改造之前先确认你的 GPU 架构到底是否支持 Tensor Core。这里也顺带提一句不少同学在 Windows 上做实验建议用 WSL2 跑 Ubuntu 或者用 Docker 镜像Windows 原生环境下的 CUDA 底层调度和显存管理经常存在一些奇奇怪怪的开销会让性能测试结果很不稳定。4. 实操过程把 AMP 移植到模型训练流程中4.1 一个可以直接抄的完整模板我以 PyTorch 官方风格为基础整理了一个最小但完整的 AMP 训练模板其中加了几个我在实践中总结过的细节配置比如 batch size 变化、learning rate 预热、多卡场景下的处理。代码结构如下import torch import torch.nn as nn import torch.optim as optim from torch.cuda.amp import autocast, GradScaler from torch.utils.data import DataLoader def train_one_epoch(model, loader, optimizer, criterion, scaler, scaler_enabledTrue): model.train() total_loss 0.0 for images, labels in loader: images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) optimizer.zero_grad() with autocast(enabledscaler_enabled): outputs model(images) loss criterion(outputs, labels) if scaler_enabled: scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() else: loss.backward() optimizer.step() total_loss loss.item() * images.size(0) return total_loss / len(loader.dataset) model create_model().cuda() optimizer optim.AdamW(model.parameters(), lr3e-4) criterion nn.CrossEntropyLoss() scaler GradScaler() for epoch in range(20): avg_loss train_one_epoch(model, train_loader, optimizer, criterion, scaler) print(fEpoch {epoch} | loss {avg_loss:.4f})为了让代码同时兼容“开 AMP”和“关 AMP”两种模式我用scaler_enabled作为开关这样在调参与对比时特别方便不用改两套脚本。使用额外配置时需要注意一下学习率调整我建议 AMP 开启后学习率策略尽量不要激进变化。虽然从理论上讲 AMP 是“无损加速”但在实际训练中高学习率窗口下GradScaler的动态调整会有短暂滞后容易造成训练初期的不稳定。可以先按原学习率跑如果发现 loss 不降再把学习率调低为原来的 0.8 到 0.9 倍。Batch size 变大后需要重新热身因为 AMP 省下了一半显存很多人顺手就把 batch size 翻倍这样确实能提高吞吐但 batch size 从 32 变到 64 后梯度噪声变了学习率不一定还能直接用。最佳做法是“batch size 翻倍学习率也跟着乘以 1.5 到 2”再进行 3 到 5 轮 warmup 观察。4.2 分布式训练里的 AMP 如何处理在单卡上AMP 的配置非常简单但在多卡DDPDistributedDataParallel环境下有一个容易踩的坑。我一开始照着单卡模式把GradScaler放在每个进程里各自创建结果发现训练 loss 不一致有时还会出现梯度不同步的问题。后来查了官方文档才发现DDP模式下GradScaler最好是每个进程创建一个并且放在DDP包装模型的外部保证梯度同步时缩放逻辑一致。更关键的是GradScaler和DistributedSampler之间没有直接关系但每个进程的数据 shuffle 必须用DistributedSampler保证这一点其实和 AMP 无关。真正与 AMP 相关的多卡配置需要把 all-reduce 操作放到“反缩放”之后PyTorch 的DDP内部已经处理好了这个顺序你只要别自己手动做梯度裁剪或梯度修改就不容易出现异常。一个省心的做法是直接用 PyTorch 官方推荐的三件套from torch.nn.parallel import DistributedDataParallel as DDP # 在初始化进程组之后 model DDP(model, device_ids[local_rank]) scaler GradScaler()记得在torch.nn.parallel.DistributedDataParallel初始化之后创建GradScaler不要提前创建否则容易遇到 CUDA context 初始化顺序不一致导致的隐性报错。4.3 推理阶段的低精度化很多人只知道训练阶段用 AMP其实推理阶段同样可以把模型转成 FP16 来提速。做法非常简单model.half() model.eval() with torch.no_grad(): inputs inputs.half() outputs model(inputs)不过要注意推理阶段并没有GradScaler所有层都直接以 FP16 运行如果你的模型里面有动态数值范围很大的操作比如某些自定义的 mask 或者 Attention Logits很容易出现输出 NaN。稳妥的做法是先用训练阶段的 AMP 精度指标做参考再在离线验证集上反复测试确认输出质量没有劣化再上生产环境。此外还有一个速度优化技巧当模型支持时可以打开torch.backends.cudnn.benchmark True让 cuDNN 在多个可选算法中自动挑一个最快的卷积实现。这个开关和 AMP 是正交关系但两者叠加起来CNN 类的训练性能还能再提升一截。5. 常见问题与排查技巧实录5.1 NaN 和 Inf最让人头疼的问题NaN 的排查几乎是混合精度训练里遇到最多的异常这里我总结了一套固定排查路径第一检查输入数据是否含 NaN。用torch.isnan(inputs).sum()或者torch.isinf(inputs).sum()快速验证。如果连输入数据都有脏值后面做啥都白搭。第二检查 loss 计算。CrossEntropyLoss 重叠了 Softmax按理说不太容易出 NaN但自定义损失函数时如果做了torch.log(0)这种操作很容易产生 -Inf再经过GradScaler放大后会更离谱。第三检查梯度。在scaler.scale(loss).backward()之后可以临时加一段代码for name, param in model.named_parameters(): if param.grad is not None and torch.isnan(param.grad).any(): print(NaN grad in, name)一旦定位到某个层的梯度出现 NaN优先检查这一层是否有exp、pow、log等容易产生大数值差异的操作。通常把这一层的输入float()保持 FP32 计算就能解决问题。第四检查 loss 本身为 0 的情况。如果GradScaler的缩放倍率被动态调整得过大极端情况下其缩放后的值也可能溢出虽然GradScaler本身有最大倍率限制但你还是可以在update()后打印scaler.get_scale()观察缩放倍率变化趋势。5.2 优化器参数和 AMP 不兼容的坑有一些优化器对 AMP 的支持并不好。比如 LAMB 优化器在大 batch 场景下对梯度的 scale 处理比较特殊如果直接搭配GradScaler使用可能造成缩放冲突。PyTorch 官方建议从 1.10 版本开始使用torch.optim.LAMB时小心处理scaler.step(optimizer)必要时可以使用scaler.unscale_(optimizer)手动反缩放再配合torch.nn.utils.clip_grad_norm_做好梯度裁剪。此外梯度裁剪的顺序也值得注意。标准流程是scaler.unscale_(optimizer) clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()如果你直接clip_grad_norm_而不先scaler.unscale_裁剪的数值实际上是在缩放后的空间里执行的数值意义不对容易把梯度裁剪阈值搞乱进而影响训练稳定性。5.3 GPU 兼容性与环境搭配问题我在操作中用到的软件组合是 Python 3.10.11 PyTorch 2.8.0 CUDA 12.1 搭配包这个组合在官方轮子下能获得不错的性能和稳定性。如果你用 Conda 管理环境建议在创建环境时提前指定 Python 版本再安装 PyTorch 对应版本这样能避免由 Python 版本不一致引起的 ABI 兼容问题。另外显存不足时直接调低 batch size 是最暴力的做法但也最省时间。很多人为了省显存把手动把某个层改成了checkpoint其实在启用 AMP 之后你可能根本不需要这么复杂的方案。反过来如果你的 PyTorch 版本比较老是 1.5 或者 1.6建议使用from torch.cuda.amp import autocast, GradScaler这种老式写法到了 PyTorch 2.x 后推荐改写成from torch.amp import autocast, GradScaler底层的精确语义没有本质变化主要是作用域更通用可以支持不同设备后端。5.4 小 batch size 场景为什么提升不明显我发现有些同学在自己实验里测不出明显的提速首先要确认一个前提batch size 是不是太小了。GPU 的利用率靠大量并行计算撑起来如果 batch size 只有 4 或者 8CUDA kernel 的启动开销和同步开销占比过高此时把计算切到 FP16虽然单次 kernel 变快了但整体等待时间并没有显著下降。这种情况下建议先把 batch size 调大或者使用梯度累积来模拟更大的 batch再去比较 AMP 的收益。梯度累积的标准写法是accumulation_steps 4 scaler.scale(loss).backward() if (step 1) % accumulation_steps 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_noneTrue)注意在梯度累积模式下scaler.step(optimizer)只能在指定时机调用。如果每一小步都调用累积的梯度就会被提前反缩放并更新掉逻辑就错了。6. 从实测数据看 AMP 的选型策略6.1 全局最优解还是局部最优解如果你同时关心训练速度和显存AMP 基本是全局最优解。但如果你需要做的是分布式训练同时又不希望额外引入复杂的第三方扩展PyTorch 自带的 AMP DDP 已经能把效率榨到不错的状态。我拿自己的实验数据再补充一个直观理解同样是跑 BERT-baseFP32 模式下显存占用 16.8GB只要我把 batch size 从 64 提到 96 就会 OOM而 AMP 模式下显存占用只有 9.1GB我可以轻松把 batch 提到 128单轮迭代吞吐量反而是原来的 2 到 3 倍。这种差距会随着模型规模的扩大被拉得越来越大。6.2 BF16 还是 FP16在大模型时代我个人的习惯是“能选 BF16 就选 BF16”。理由很简单BF16 的动态范围和 FP32 完全一致训练过程中的损失曲线更稳几乎不会出现下溢导致“死训练”的问题。但 BF16 也有一个明显的短板就是低精度时能保留的有效位很少对极其敏感的小数值更新不如 FP16。不过由于 PyTorch 的混合模式会在关键节点自动切回 FP32所以这个问题在工程实践中影响有限。如果你的卡是 RTX 30 系列或更新的 Ampere 架构直接用 BF16 很省心。如果你是 RTX 20 系列建议老老实实用 FP16配合GradScaler同样可以获得非常可观的加速。6.3 混合精度和模型并行、流水线并行的协同在大模型训练中AMP 通常还会和模型并行、流水线并行一起使用。比如你在用 DeepSpeed 的时候它会自动在 ZeRO 优化器和 AMP 之间做一体化配置。这时候要小心DeepSpeed 自带了一套混合精度的封装fp16配置选项和 PyTorch 的GradScaler不能同时开否则会发生双重缩放导致 loss 变成天文数字或者 NaN。如果项目是从零搭建建议先只开 PyTorch 原生的 AMP再加 DDP跑通并确认精度和速度之后再决定是否引入 DeepSpeed 等框架。一步步来别一次性叠太多优化策略不然出了问题你连源码都不太想翻。7. 训练效果与扩展方向的实战体验7.1 图像分类任务中的实际表现我拿 ImageNet 子集做过完整实验ResNet-50 在开启 AMP 后每个 epoch 的训练时间减少了 35% 到 40%而 Top-1 准确率和 FP32 版本基本是持平的。这个结论在多个标准 CNN 上都能复现这说明对于“数值分布相对平稳”的卷积网络AMP 几乎是无损的。7.2 目标检测与分割任务里的技巧在检测与分割任务中因为损失函数往往包含多个子项分类损失、回归损失、mask 损失不同子项对精度的敏感度不同所以你要注意把汇总后的 loss 交给scaler.scale()而不是对每个子项分别 scale。例如loss loss_cls loss_box loss_mask这样做的原因是GradScaler会对整体 loss 做一次统一缩放如果每个子 loss 分别缩放再相加整个计算图的对齐关系就会乱套梯度在反向传播时无法正确汇聚。7.3 文本生成与 LLM 微调场景文本生成模型GPT 系列等在微调时AMP 带来的收益也很大尤其在使用 LoRA 这类参数高效微调时。不过需要注意的是LoRA 和 AMP 的配合稍微有点讲究LoRA 的权重通常以 FP32 存储但 forward 阶段经过autocast会变成 FP16这在多数情况下没问题。但如果你把 LoRA 的 rank 调得特别大或者对精度敏感的任务建议在测试阶段对比 FP32 和 AMP 的验证集指标别盲目信任速度提升。我还见过一种做法是在训练过程中设计“阶段式精度切换”前几个 epoch 用 FP32 稳定收敛中间切到 AMP 提速最后几个 epoch 再切回 FP32 微调。这种方式本身有其价值但在常规任务中并不需要只在超大规模模型或精度要求极其严格时值得尝试。8. 关于 PyTorch 生态与社区实践的一些心得8.1 PyTorch 版本和生态对 AMP 的影响PyTorch 的 AMP 实现已经非常稳定从 1.6 到 2.x核心 API 基本没变过。如果一定要选版本我建议用 PyTorch 2.0 以上的版本因为其内部对 CUDA graph、torch.compile 的支持更完善。torch.compile开启后可以把 AMP 的算子融合得更深进一步减少 kernel 启动开销实测性能还能再提升 10% 到 20%。不过torch.compile的编译时间比较长第一次跑会明显卡顿后续会走缓存速度正常。8.2 社区配套工具和资源目前 PyTorch 的 AMP 功能已经内置于torch.cuda.amp和torch.amp并不需要额外安装第三方包。但如果你在 HuggingFacetransformers框架里做模型微调它自带的Trainer支持在TrainingArguments中直接传fp16True或bf16True底层就是调用了 PyTorch 的 AMP 功能非常省事。我自己在做文本模型实验时很少再去手写训练循环直接用 Trainer 就能把混合精度跑起来。8.3 版本兼容与依赖安装的坑Python 版本和 PyTorch 版本的搭配是整个环境的基础。我在自己的项目环境里推荐 Python 3.10因为目前主流库对 3.10 的兼容性最好。安装 PyTorch 时优先选择官方渠道提供的 CUDA 版本 wheel不要图省事选用别人整合的“一键包”避免 CMS 底层库冲突。把这些都整理好之后AMP 的使用几乎没有额外负担。9. 最后的避坑清单为了让这篇文章的实操价值更高我把容易忽略的要点整理成一个速查清单每一条都是我踩过坑之后沉淀下来的。检查 GPU 是否支持 Tensor CoreAmpere 及以上最好否则 AMP 提升可能有限。建立基线改造前先跑纯 FP32 的数据记录耗时、显存、loss 曲线和验证集指标后面互相参照才能定位问题。训练循环里永远使用with autocast():包裹 forward 和 loss 计算不要手动把模型整体half()。GradScaler必须搭配使用不建议单独使用autocast。梯度裁剪之前先调用scaler.unscale_(optimizer)。多卡训练时把GradScaler放在 DDP 包装之后创建。在 checkpoint 保存时建议把scaler.state_dict()一并保存恢复训练时再加载否则恢复训练时缩放倍率需要重新适应容易造成训练 spike。加载保存的示例checkpoint {model: model.state_dict(), optimizer: optimizer.state_dict(), scaler: scaler.state_dict()} torch.save(checkpoint, ckpt.pth)恢复时checkpoint torch.load(ckpt.pth) model.load_state_dict(checkpoint[model]) optimizer.load_state_dict(checkpoint[optimizer]) scaler.load_state_dict(checkpoint[scaler])我之前忘了保存scaler.state_dict()在长周期训练中断后恢复结果重启后前几个 step 的 loss 异常剧烈波动当时排查了很久最后才发现是这个细节造成的。这种小坑书面文档里很少会特地强调但实际项目里真的会咬人。还有一点AMP 的收益和模型类型密切相关。像 CNN、Transformer、Embedding 密集型模型AMP 收益大而像某些数值计算密集、动态范围特别敏感的模型比如涉及大量自定义 CUDA kernel 的操作AMP 可能会带来不可预期的精度损失。遇到这种情况建议把不稳定的子模块排除在autocast之外用autocast(enabledFalse)包裹这部分前向计算或者专门写一个 CUDA 扩展来保持精度。10. 关于下一步优化的方向AMP 只是优化训练的第一步。当你已经用上了 AMP接下来可以考虑的方向有这么几个打开torch.compile融合算子和减少内存拷贝使用channels_last的显存布局在 CNN 场景里能搭配 Tensor Core 做更好的性能调度使用DataLoader的pin_memoryTrue和non_blockingTrue把数据加载的时间压下去如果显存还是不够梯度 checkpoint 加 AMP 双管齐下如果训练集群本身有多个节点把 AMP FSDPFully Sharded Data Parallel组合使用能更极限地优化显存和训练吞吐。我个人在实际操作中的一个体会是AMP 的收益并不是靠某个“高深莫测”的配置实现的它其实就是把硬件本身的能力用起来了。很多同学卡在性能上不是因为模型写错了而是因为默认 FP32 这个习惯限制了 GPU 的发挥。从 FP32 切换到混合精度通常只需要改十几行代码换来的是训练时间缩短三分之一、显存砍半这种投入产出比在深度学习优化里真的不多见。最后再分享一个小技巧如果你不确定自己的模型经过 AMP 后精度是否受了影响可以取训练过程中的一个 checkpoint分别在 FP32 和 AMP 模式下做相同数量的前向推理对比模型输出的数值相似度。torch.max(torch.abs(fp32_out - amp_out))如果小于 1e-2 量级基本可以放心继续使用。这个方法不需要重新训练五分钟就能定位问题很实用。混合精度训练这条路走到今天已经是深度学习的标配技能。无论你是刚入门还是已经在做大规模训练掌握好autocast和GradScaler把 AMP 合理嵌入训练流程都会让模型的迭代速度和工作效率明显上一个台阶。希望这篇总结能帮你少走一些弯路。
返回列表