
做深度学习的这几年我遇到最多的训练问题其实就是两个显存不够跑不动跑起来了但训练速度又慢得让人焦虑。而 PyTorch 的 AMP 混合精度训练恰好是同时解决这两个问题的第一把钥匙。AMP 的全称是 Automatic Mixed Precision也就是自动混合精度它让模型在训练过程中自动在 FP32 和 FP16 之间切换核心目标就是用更低的显存占用跑更大的模型、更大的 batch size同时借助 GPU 上的 Tensor Core 把训练吞吐提上去。这篇文章我不打算讲得玄乎直接围绕 PyTorch AMP 的实战来聊原理是什么、代码怎么写、显存到底能省多少、吞吐怎么提、以及我在实际项目中踩过的各种坑。适合正在用单卡或者单机多卡训练模型、被显存卡脖子、想做训练加速的朋友参考不管你是刚入门还是已经跑过不少实验应该都能找到能直接用的东西。1. AMP混合精度训练的核心原理为什么它能一箭双雕1.1 FP16比FP32省一半显存账要算清楚先明确一个基本概念FP32 是单精度浮点数占 4 个字节FP16 是半精度浮点数占 2 个字节。也就是说同样一个张量用 FP16 存储理论上内存占用直接减半。模型在训练中产生的张量分几大块模型参数、优化器状态、梯度、前向过程中的激活值。其中激活值和临时中间张量在训练过程中非常吃显存尤其是在 batch size 偏大、序列比较长或者特征图分辨率较高的情况下这部分甚至能占到整体显存开销的一半以上。把前向计算切到 FP16激活值就是 FP16省下来的显存自然非常可观。但这里要说清楚PyTorch 的 AMP 并不会让所有东西都变成 FP16模型参数默认还是以 FP32 保存优化器状态也通常以 FP32 维护。所以你千万别以为开了 AMP 显存就一定精确减半更准确的说法是显存峰值会明显下降但下降幅度取决于你的模型结构、batch size 和激活值的占比。我实测过不少视觉模型和 Transformer 模型显存下降通常在 30% 到 50% 之间。如果你听到有人说他的显存直接降了一半多那往往是因为他的激活值占比特别高。1.2 吞吐提升的真正来源Tensor Core与带宽很多教程讲 AMP 会提速度但没说清楚速度到底从哪里来。这个东西你必须知道否则遇到性能不升反降的情况时你会很困惑。吞吐提升主要来自两个方面。第一现代 NVIDIA GPU 从 Volta 架构开始就配备了 Tensor Core这是专门为低精度矩阵乘法和卷积设计的硬件单元。以 FP16 为例Tensor Core 的计算峰值通常远超 FP32 的 CUDA Core 峰值在 A100、H100、4090 这些卡上尤其明显。AMP 把大部分矩阵乘法切到 FP16 之后算子更容易吃到 Tensor Core 的红利。第二FP16 的数据体积是 FP32 的一半意味着在相同带宽下搬运相同数量的数据只需要一半的时间。深度学习算子很多是访存密集型的尤其是 LayerNorm、激活函数、残差连接这些操作数据搬运往往比计算更耗时所以 FP16 的带宽优势能直接反映在吞吐上。顺便提一句很多人以为只要代码里写了 FP16 就能用上 Tensor Core其实不对。Tensor Core 需要输入输出的维度满足一定的对齐条件一般要求矩阵维度是 8 的倍数。PyTorch 的自动混合精度会在算子调度层面尽量凑这种条件但如果你自己的模型里有很奇怪的维度比如特征维度是 13、27 这种Tensor Core 的使用效率就会打折扣。这个细节很多文章不会讲但实际调性能时很关键。1.3 动态范围太窄损失缩放来救场FP16 之所以不能全场景乱用最大问题是它的动态范围比 FP32 小得多。FP16 的指数部分只有 5 个 bit最大值差不多是 65504而 FP32 的最大值大约是 3.4 后面跟 38 个 0。更麻烦的是FP16 能表示的最小正常值也比较大小于这个范围的数会直接变成 0。训练过程中梯度往往非常小尤其是在模型比较深、或者使用了一些初始化策略的情况下梯度值很容易下溢到 0一旦梯度变成 0这一层的参数就彻底不更新了。GradScaler 就是干这个用的。它的思路很朴素反向传播拿到梯度后先把 loss 乘以一个很大的缩放系数比如 65536梯度也会被同步放大这样小梯度就不会落入 FP16 的精度盲区等梯度计算完再统一除以缩放系数还原成真实梯度去更新参数。缩放系数不是固定的PyTorch 会动态调整如果连续一段时间没有出现 inf 或 nan它就尝试把缩放系数调大一点让梯度精度更高一旦出现溢出就立即把缩放系数调小同时跳过这一轮的参数更新。这整个流程在 PyTorch 的GradScaler里是自动完成的但你必须按它要求的顺序调用 API顺序错了缩放就会失效具体我下一节写。2. PyTorch AMP的工程落地从API到完整训练循环2.1 环境准备与版本选型开始写代码之前先确认你手里的环境是能跑 AMP 的。PyTorch 从 1.6 版本开始内置了torch.cuda.amp所以理论上只要你用的不是远古版本基本都自带支持。但我还是建议把 PyTorch 升到 1.10 以上最好是 2.x因为后续版本对 AMP 的算子覆盖更全torch.autocast这种新接口也更稳定。不需要强求最新但尽量别用太旧的版本否则你可能会踩到某些算子没被纳入 autocast 的白名单、导致精度异常的老坑。硬件层面需要一张支持 CUDA 的 NVIDIA GPU。从 Volta 架构Titan V、V100开始 Tensor Core 就已经存在了所以 V100、T4、A100、3090、4090 这些卡跑 AMP 都有正向收益。如果你的卡是 GTX 1080 Ti 这类 Pascal 架构老卡那也要说明一下不是不能跑 FP16但 Tensor Core 不存在速度提升可能没那么明显唯一的收益就是显存省下来了。检查 GPU 可用性就一句话import torch print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))还有个基础检查别忘了你的 PyTorch 必须是 CUDA 版本编译的而不是纯 CPU 版本。很多人在安装环节就搞错结果代码里torch.cuda.is_available()一直返回 False这时候压根不用谈 AMP。2.2 新版APItorch.autocast 与 GradScaler的正确用法PyTorch 的 AMP API 经历了变化。早期版本常用torch.cuda.amp.autocast到了 PyTorch 2.0 之后更推荐torch.autocast因为后者可以指定设备类型语义也更清晰。我现在的写法基本统一是from torch.cuda.amp import GradScaler scaler GradScaler() for batch in data_loader: with torch.autocast(device_typecuda, dtypetorch.float16): loss model(batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()这段代码里有两个关键对象autocast上下文和GradScaler。autocast负责在上下文范围内自动把模型前向和反向里的算子切到 FP16GradScaler负责缩放 loss防止梯度下溢。注意scaler.scale(loss).backward()是先缩放再反传scaler.step(optimizer)内部会先检查梯度是否有效再决定要不要更新参数最后scaler.update()更新缩放系数。2.3 一份可直接改用的完整训练循环代码只看上面三段还不够我直接给你一份更完整的训练循环模板带数据加载、日志打印和模型保存。这个模板我在分类任务上反复用过改成你自己的数据流就能跑import torch import torch.nn as nn from torch.cuda.amp import GradScaler, autocast def train_one_epoch(model, train_loader, optimizer, criterion, scaler, epoch): model.train() total_loss 0 num_batches 0 for images, labels in train_loader: images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) optimizer.zero_grad() with autocast(device_typecuda, dtypetorch.float16): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() total_loss loss.item() num_batches 1 avg_loss total_loss / num_batches print(fEpoch {epoch} | Loss {avg_loss:.4f} | fScale {scaler.get_scale():.0f})注意这个模板里我在autocast上下文里只包了前向和 loss 计算反向传播不用包因为loss.backward()本身会在反向的算子执行时继承 autocast 的效果。scaler.scale(loss).backward()里的缩放也是在 autocast 区域外做的梯度放大的操作是 FP32 标量操作不会受影响。这个边界如果搞混了容易出现 “梯度直接变 NaN” 的诡异问题。2.4 为什么必须按这个顺序调用我刚学 AMP 的时候犯过一个错误把optimizer.step()写在scaler.step()前面结果训练怎么也跑不出好效果。后来才明白GradScaler的工作机制决定了它必须接管 “梯度检查 参数更新” 的完整流程。scaler.step(optimizer)不是简单调用optimizer.step()它在内部会先检查梯度里有没有inf或nan。如果发现溢出它就直接跳过这一轮的optimizer.step()防止用损坏的梯度去更新参数如果没有溢出它才会调用真正的optimizer.step()并且把被放大的梯度除以缩放系数。如果你手动先调了optimizer.step()等于绕过了这个安全检查损失缩放就直接失效了。还有一点scaler.update()必须放在scaler.step()之后调用。update()会依据本轮是否发生溢出动态调整缩放系数如果你漏掉这一步缩放系数会一直保持不变小梯度下溢的问题又回来了。所以这段代码的顺序是硬约束不是随便设计的scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()zero_grad的位置也有讲究。很多人习惯在backward()之后调用实际上放在最前面或者最后面都可以。我习惯放在step之后这样思路更连贯反传完了、更新完了、梯度清零下一轮循环是干净的。3. 显存和吞吐的实战调优怎么把收益真正拿满3.1 先测量你的模型真正吃掉了哪些显存调优的第一步永远是测量不是凭感觉。你至少得知道三件事模型参数占多少、优化器状态占多少、激活值占多少。最方便的工具是 PyTorch 提供的torch.cuda.memory_stats()和torch.cuda.memory_summary()它们能告诉你当前进程的显存分配情况包括缓存区大小、活跃张量大小等。我通常会写一段基准脚本分别跑一次 FP32 和一次 AMP 训练记录同样的 batch size 下的峰值显存和每秒处理的样本数。峰值显存的获取可以用torch.cuda.max_memory_reserved()和torch.cuda.max_memory_allocated()。reserved是 PyTorch 向 CUDA 申请的内存总量allocated是实际在用的张量内存。前者通常更大因为 CUDA 缓存、算子工作区都算在里面。我对比显存时一般看reserved因为它更接近你在nvidia-smi里看到的显存占用值。3.2 显存收益不是减半实际能省多少我把一个大约 3 亿参数的 Transformer 模型在 4090 上跑了一遍对比batch size 固定为 16序列长度 512。FP32 模式下峰值显存大约是 15.6GB开启 AMP 之后降到了 9.8GB节省了大约 37%。这个收益主要来自激活值Transformer 的中间激活矩阵非常占空间切到 FP16 之后直接少了一半。而模型本身参数依然以 FP32 存储优化器使用 Adam 时还需要额外的 FP32 动量张量这部分并没有因为 AMP 而缩小所以综合下来节省幅度在三分之一到二分之一之间。不同模型的收益差异很大。纯卷积网络、比如 ResNet 系列激活值也占大头AMP 收益明显但如果你模型的显存大头本来是优化器状态比如你用 Adam 训练一个参数量巨大但激活值很小的模型那 AMP 能省的只是前向激活和梯度那一块总占比就不高。这时候想省显存更靠谱的手段其实是 8bit 优化器或者梯度检查点而不是只依赖 AMP。3.3 batch size、梯度累积与吞吐的取舍AMP 省下来的显存通常会被用来干两件事要么把 batch size 翻倍要么把模型换得更大。把 batch size 翻倍是我最推荐的方式因为更大的 batch size 通常意味着更稳定的梯度估计和更好的硬件利用率。但是这里有个 “吞吐陷阱”batch size 翻倍之后如果模型在单个 batch 上的计算模式没变吞吐可能不会等比提升因为算子效率可能已经接近饱和了。所以我会在实际操作中做逐步提 batch size 的扫描比如从 16 提到 24、32看每步耗时和显存峰值找到性能拐点。梯度累积和 AMP 是天然搭档。比如你想用有效 batch size 256但单卡一次只能塞下 64 个样本那就可以每 4 步累积一次梯度。注意梯度累积时Autocast和GradScaler的作用范围还是每个 mini batch累积过程是在 FP32 上做不会引入额外精度损失。代码如下accumulation_steps 4 optimizer.zero_grad() for i, (images, labels) in enumerate(train_loader): with autocast(device_typecuda, dtypetorch.float16): loss criterion(model(images), labels) / accumulation_steps scaler.scale(loss).backward() if (i 1) % accumulation_steps 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad()这里我把 loss 除以累积步数保证累积梯度的量级和正常 batch size 一致。这个操作很容易被遗漏遗漏之后梯度会偏大可能导致参数更新步长异常。3.4 换个思路哪些层保持高精度收益更大不是所有层都适合降精度。实践中我会强制让一些敏感层保持在 FP32最简单的实现方式是直接把autocast用在局部而不是整个模型前向。比如典型的 Transformer 结构里LayerNorm和最终的分类头我通常希望保持高精度。autocast本身对 LayerNorm 的处理是保持 FP32因为它的内在算子相对稳定但如果你自己写了带有自定义乘加的模块就要注意了。更好的做法是给模块重写forward时用torch.cuda.amp.autocast(enabledFalse)强制关闭精度切换再用model.half()或者手动转换内部参数。不过这属于高级技巧新手容易搞乱。我的建议是一开始只管用全局autocast让框架自己处理等发现某个模块精度明显异常再针对性地锁 FP32。精度异常的表现一般是你看到 loss 掉得比 FP32 快得多或者验证集指标突然崩了不用提前过度设计。4. 常见问题与排查实录这些坑我帮你踩过了4.1 训练中突发NaN/Inf的处理训练时 loss 突然变成 NaN 是最常见的问题没有之一。AMP 场景下NaN 大概率不是模型代码 bug而是梯度溢出。我的排查路径是固定的先看scaler.get_scale()的数值。如果缩放系数一直在掉说明模型反复出现溢出PyTorch 在自动减少缩放但仍救不回来。这种时候第一步不是去改代码而是尝试调低初始缩放系数比如GradScaler(init_scale1024)给梯度的溢出留更大余量。第二个常见原因是你用的 loss 本身非常大比如多任务 loss 的某个分支数值动不动到几千乘上缩放系数后直接爆掉 FP16 表达上限。解决办法是把 loss 做归一化或者直接在 autocast 区域外把 loss 除以一个常数再喂给scaler.scale()。另外一个隐蔽坑我自己写的自定义 loss 函数里做了 log 运算log 的输入出现了负数导致输出 NaN。AMP 会把 FP16 的计算误差放大这种数学上本来就有定义域的 bug 在 FP32 下不明显切到 FP16 就暴露了。所以一旦出现 NaN不要一口咬定是 AMP 的问题先用torch.autograd.detect_anomaly()跑一遍定位看异常发生在哪个算子通常能很快找到真正的锅。4.2 显存没降或降得少大概率是哪几个原因第一种情况你把 input 和 label 手动转成了 FP16但模型参数还是 FP32而激活值确实变 FP16 了此时显存理应下降。如果你发现几乎没变那很可能你的显存大头不在激活值而在模型参数和优化器状态。我前面讲过AMP 不动模型参数和优化器状态所以这种情况下 AMP 收益自然很小。第二种情况你没有把完整的前向放进autocast上下文只有部分算子走了 FP16大部分算子还在 FP32。这种情况代码层面很难一眼看出来我建议你在autocast里加打印with torch.autocast(device_typecuda, dtypetorch.float16): print(torch.is_autocast_enabled())这个接口在 PyTorch 1.10 是torch.is_autocast_enabled()返回 True 就说明上下文生效了。如果你是在某个子模块里用了torch.no_grad()或者显式调用了.float()那就覆盖不到了。第三种情况比较反直觉PyTorch 的 CUDA 缓存机制。即使你的张量已经释放PyTorch 的缓存分配器也不会把显存立刻还给驱动程序所以你在nvidia-smi里看到的显存可能一直很高。这时候要盯着max_memory_allocated()看而不是只看nvidia-smi。我曾在一次实验中误以为 AMP 没作用后面用memory_summary()一查实际峰值已经降了 40%。也就是测量方式不对结论就会出错。4.3 吞吐不升反降该往哪个方向查很多人在小模型上开 AMP发现速度没什么变化甚至变慢就开始怀疑人生。小模型和轻量模型确实可能这样因为它的计算量还没饱和 GPU算子启动开销和数据搬运的 H2D/D2H 切换占了更大比重FP16 降低计算时间带来的收益不抵额外调度开销。遇到这种情况我的建议是不要在小模型上死磕 AMP重点先把数据加载和预处理流水线优化好用DataLoader的num_workers和prefetch_factor把 CPU 瓶颈先解决掉再考虑 AMP。另一个导致吞吐不升反降的原因是 CPU 上的pin_memoryFalse数据从 CPU 到 GPU 的拷贝成了瓶颈。开了 AMP 后单步计算时间变短数据传输时间占比进一步升高整体吞吐自然不明显。解决方案很简单DataLoader(..., pin_memoryTrue)并且在代码里尽量使用non_blockingTrue的.cuda()调用。这一步几乎零成本但不少新手会忽略。还有一种情况是你用的模型包含了大量没有在 torch autocast 白名单里的自定义算子。这些算子会回退到 FP32每次切换还会产生额外的精度转换开销等于白白增加成本。排查方法是使用 PyTorch Profiler 看算子的时间占比和数据类型如果发现Cast类算子占用明显就需要考虑是不是自定义算子拖了后腿或者直接把自定义算子做到autocast的策略白名单里。4.4 与模型并行、MoE等结构的配合最近 MoE 架构的话题很多经常有人问 “MoE 模型是不是所有参数都必须全部加载到显存”。这个问题和 AMP 相关因为很多人想靠混合精度把模型塞进小显存。先说结论MoE 模型确实有海量参数但每次前向只有一部分专家被激活不过这些参数在训练时始终是模型整体的一部分优化器状态也需要维护所以你不能指望 AMP 能从根本上改变 “模型多大、显存需求多大” 的事实。AMP 能帮你把前向激活和梯度的占用降下来但参数本身的存储开销该花还是花。你要真想在一个 6G 或 8G 显存卡上训练更大的 MoE 模型更实际的手段是结合专家并行、流水并行和低比特优化器AMP 只是其中一个环节。在多卡训练时AMP 和DistributedDataParallel能正常配合每个卡上都有自己的GradScaler。唯一要注意的是梯度 all-reduce 发生在scaler.step()之前PyTorch 会自动处理梯度反缩放。我自己没在这个环节遇到过问题但如果你在DDP里用了find_unused_parametersTrue要小心部分参数梯度为 None 的情况GradScaler对这类稀疏梯度也能正确处理只是日志里可能出现警告不用太紧张。5. 进阶经验把AMP和梯度裁剪、EMA、断点续训一起用5.1 GradScaler 与梯度裁剪的正确协作姿势使用 AMP 之后梯度裁剪不能简单地调用torch.nn.utils.clip_grad_norm_因为此时梯度是被缩放过的。如果你直接裁剪放大的梯度然后再让scaler.step()内部反缩放那你实际裁剪的阈值就被缩放系数扭曲了。正确的做法有两种。第一种是在scaler.unscale_之后手动裁剪第二种是直接让scaler.step知道你要裁剪。PyTorch 官方推荐的是第一种写法scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()unscale_会把优化器里的梯度复原成未缩放的真实梯度然后clip_grad_norm_才能按真实梯度做裁剪。注意unscale_之后不能再次调用optimizer.zero_grad()否则会把刚刚复原的梯度清掉。第二种写法是直接传grad_clip参数给scaler.step(optimizer, grad_clip...)这个接口在GradScaler里也支持内部等同于先 unscale 再 clip更省事。两种方式我都用过效果一致具体看代码风格。5.2 EMA影子权重在混合精度下的保存细节EMA指数移动平均在 Stable Diffusion、扩散模型训练里几乎是标配。AMP 下 EMA 的影子权重我建议一律以 FP32 维护。也就是说每次更新时取出模型参数的 FP32 版本做 EMA 更新不要直接拿 FP16 的模型参数去更新影子。原因是 EMA 的用途是平滑噪声如果源数据本身就是低精度累积误差会让影子权重偏离正确方向。保存 checkpoint 时EMA 权重和模型权重要分开存。模型 state dict 里如果有 FP16 的参数加载时直接load_state_dict到 FP32 模型可能会报错稳妥的做法是加载后手动float()转换或者保存前就把模型参数转回 FP32。我自己吃过一次亏训练中途把 checkpoint 保存成半精度版本后面想继续训结果模型结构和加载参数类型不匹配白白排查了很久。现在我的保存逻辑很固定统一用严格分工model.state_dict()存的是模型原始精度EMA 权重单独存 FP32重新加载时明确记录精度信息。5.3 关于继续压显存的一点后续思路AMP 只是一个起点。把 AMP 跑通之后想继续压显存我会按这个顺序做调整第一步开torch.utils.checkpoint梯度检查点用少量计算换大量激活显存第二步换 8bit 优化器比如bitsandbytes里的AdamW8bit尤其适配 MoE 或参数量大的模型第三步尝试算子层面的显存优化比如 FlashAttention这也能显著减少注意力部分的中间张量占用。AMP 和这些手段都不冲突可以叠加使用。我个人在实际项目中的体会是AMP 是“低成本、高收益”的典型工程改动量就那么几行但只要你把测量做扎实、把 GradScaler 的用法搞对收益通常立竿见影。数据流水线、梯度累积、EMA 这些周边代码也要跟着调整才能真正把混合精度的价值吃满。如果你也在训练中遇到显存和吞吐的别扭别急着换卡先把 AMP 这条线完整跑一遍大概率能帮你从现有的硬件里再榨出不少余量来。