PyTorch AMP混合精度训练实战:显存优化、Loss Scaling与踩坑指南
前阵子用一张24G的卡训练一个1.4B的稠密模型batch size调到2就直接OOM。排查来排查去静态内存算下来其实根本没超卡死就死在激活值这一块。后来把AMP接进去峰值显存直接掉了差不多三分之一吞吐还反向涨了30%多。这篇文章就把PyTorch AMP混合精度训练这件事彻底讲透从显存账单怎么算、AMP为什么能省又为什么不能什么都省到实际接入的完整代码、loss scaling机制、多卡场景的注意事项以及我踩过的几个一眼难看的坑全给你盘一遍。混合精度训练不是什么黑魔法它的核心就一句话能在半精度FP16/BF16下算的算子尽量用半精度算必须保持精度的部分继续用FP32。但这背后的显存分配逻辑、梯度缩放原理、以及不同卡型、不同模型结构下的实际收益差距如果不搞清楚很容易出现加了AMP反而更慢loss突然变NaN显存没降多少这类莫名问题。1. PyTorch训练时的显存账单先把钱花在哪弄明白1.1 训练显存四大开销参数、梯度、优化器状态、激活值任何训练任务在显存里占地方的无非四类模型参数、梯度、优化器状态和前向过程中的中间激活值。很多人只算第一项模型参数4字节一个然后被OOM打脸就是因为把后三样当空气了。以参数量为P单位B十亿参数的模型为例FP32训练时模型参数P × 4GB。1B就是4GB。梯度和参数一样大再看4GB。Adam优化器状态每个参数要存一阶动量m和二阶动量v都是FP32所以是 P × 8GB。1B就是8GB。激活值这个没法简单公式化和batch size、序列长度、网络深度直接相关。但大模型训练时激活值动辄就是几GB到几十GB。光是前三项1B模型用AdamW就要吃掉16GB。你把batch调大一点激活值冲上来24G卡必炸无疑。这不叫显存不够这叫账单没算清。很多人问大模型总参数和激活参数对显存的影响说的就是这个——总参数决定静态占用激活参数决定动态峰值两者叠加才是真实的峰值显存需求。1.2 为什么实际显存永远比你算出来的大CUDA context和allocator多出来的显存去哪了我自己实测的经验光是一次CUDA初始化context、cuDNN、cuBLAS这些库就要占掉500MB到1GB不等你一个模型还没建出来这笔钱就花出去了。PyTorch的显存分配器会按需向CUDA申请显存但释放的时候不是立刻还给GPU而是留在自己的缓存池里方便下次复用。这导致nvidia-smi看到的占用会比PyTorch财报torch.cuda.memory_summary()高出一截。跑长训练任务时如果反复创建不同shape的张量还可能出现显存碎片——剩余空间很多但找不到连续的一大块。所以看显存占用训练过程中要用PyTorch自己的统计import torch # 每隔一段时间打印一次峰值显存 print(torch.cuda.max_memory_allocated() / 1024**3, GB) print(torch.cuda.max_memory_reserved() / 1024**3, GB)前者是实际分配出去的张量总量后者是PyTorch向CUDA申请的总量。两者差出来的就是allocator的缓存池和碎片。我的习惯是保留一条log定期输出这两个值配合nvidia-smi看才能真正定位问题是在模型本身还是分配策略上。1.3 AMP到底打在了账单的哪一行关键点来了AMP并不会把模型参数、梯度、Adam状态这些主账目变成FP16省一半。AMP的做法是模型参数主副本仍然是FP32优化器状态也仍然是FP32。它真正省的是激活值和计算过程中的临时张量——这部分在autocast上下文里自动被降成FP16直接腰斩。同时反向传播时有些梯度中间结果也会以FP16形式存在。这就是为什么很多人开了AMP后显存只降了20%-40%而不是50%。如果你的模型本身就小、batch又不大激活值占比不高那你开AMP会发现显存掉得不多——因为省钱的地方本来就没花多少钱。反过来大模型、长序列、大batch的场景激活值动辄占掉显存一半甚至更多AMP的收益就非常可观。2. AMP的核心机制FP16、BF16、Tensor Core和loss scaling2.1 FP16为什么快、为什么危险FP16用1位符号、5位指数、10位尾数表示数。它的最大有限值是65504这不算大更麻烦的是精度——相对精度大约只有FP32的千分之一左右。这意味着两件事数值太大超过65504会变成inf这是FP16的溢出问题overflow。数值太小低于大约5.96e-8的正数会变成0这是下溢问题underflow对训练更致命。训练中梯度经常是小数值比如1e-10这种量级。如果梯度被存成FP16再经过backward直接就是0这个参数等于没梯度。AI框架常规FP16训练会用一个loss scale来解决下溢在forward算出的loss上乘一个大数常用2^16也就是65536整体数值都被放大backward得到的梯度同样被放大就不会掉到FP16的最小表示范围以下了。反传结束后再把梯度缩小回去让优化器在FP32的精度下正常更新。2.2 GradScaler和autocast各管哪一段PyTorch的AMP由两个组件配合autocast负责在forward中自动把某些算子切换为FP16执行GradScaler负责ls的放大、缩小以及防止NaN/Inf的溢出保护。这段是很多教程没说透的autocast不是把所有层全部切成FP16而是按算子白名单来。比如矩阵乘法、卷积、linear这些计算密集又对精度损失不敏感的算子用FP16而softmax、layernorm、loss计算这些需要动态范围或者数值稳定的算子即使写在autocast里也仍然用FP32计算。你不用自己判断哪些能转哪些不能转PyTorch内部已经帮你分好类了。那为什么还要GradScaler因为autocast只管精度切不切换它管不了梯度下溢这件事。你需要把loss放大N倍再在反传后缩小梯度。这套逻辑全在scaler里。scaler.step(optimizer)的动作也不是简单调一下optimizer.step它内部会检查这次迭代的梯度里有没有inf或NaN。如果有它会跳过这次的参数更新并降低scale如果连续一段时间的迭代都正常它会尝试调大scale让梯度尽量落在FP16能表示的范围里。2.3 BF16另一种半精度大模型训练的新宠FP16有动态范围窄的问题BF16直接把指数位扩成8位和FP32一致所以它的动态范围跟FP32差不多不会轻易inf或下溢。代价是尾数只有7位精度更低。BF16的好处是不需要loss scaling这套机制也不用担心缩放后的梯度归零。坏处是精度更低所以目前常见用法是大规模预训练、微调场景用BF16为主FP16更多用在检测、分类这类相对温和的下游任务上。HMEL和A100以上显卡对BF16的Tensor Core支持都很好。如果你的卡支持BF16RTX 3090/4090、A100、H100这些大模型我推荐优先试BF16。PyTorch里只需要把dtype参数换成torch.bfloat16autocast的使用方式一模一样。2.4 Tensor Core为什么算得快不是玄学AMP能提升训练吞吐除了显存带宽压力小了之外核心硬件原因是NVIDIA的Tensor Core。Tensor Core本质上是一个专门为半精度矩阵乘法设计的加速单元它的 FP16/BF16 算力通常是FP32的两倍甚至更多。你的矩阵乘法和卷积一旦落到Tensor Core上吞吐自然上去。不过这有前提你的模型得是计算密集型。如果模型本身小、并发少、或者数据IO瓶颈远大于计算瓶颈那不管什么精度都救不了你——这解释了为什么有些小模型开AMP之后速度没变化甚至因为精度转换、额外logits检查反而变慢了一点。3. 实战接入给训练脚本装上AMP的标准姿势3.1 一个可以直接抄的最小模板假设你原先的训练循环长这样for x, y in dataloader: x, y x.cuda(), y.cuda() optimizer.zero_grad() loss model(x, y) loss.backward() optimizer.step()接入AMP只需要三步初始化GradScaler、用autocast包forwardloss、把backward和optimizer.step换成scaler的接口。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for x, y in dataloader: x, y x.cuda(), y.cuda() optimizer.zero_grad() with autocast(): output model(x) loss loss_fn(output, y) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()就这六行改动。scaler.scale(loss).backward()是先放大loss再反传scaler.step(optimizer)是检查梯度情况后决定是否真正更新参数scaler.update()负责动态调整scale值。注意autocast作用范围应该包含loss计算因为在计算交叉熵这类loss时如果对logits做FP16输入CUDA算出来的分布会更贴FP16而loss本身仍然是FP32。3.2 梯度累积时的一个坑如果你做梯度累积最省事的写法是等累积步数到了再调scaler.step和scaler.updateaccumulation_steps 4 for i, (x, y) in enumerate(dataloader): x, y x.cuda(), y.cuda() with autocast(): loss model(x, y) / accumulation_steps scaler.scale(loss).backward() if (i 1) % accumulation_steps 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad()这里有个细节因为每次backward都会把梯度放大同样的scale值累积之后再一起step整体scale是统一且一致的没问题。但如果中途又调用scaler.unscale_去手动缩放梯度那一定要保证每个step只调用一次否则梯度会被缩放两次直接把数值搞崩。3.3 梯度裁剪的顺序先unscale再clip这是AMP最常见的翻车点。torch.nn.utils.clip_grad_norm_的输入梯度得是真实幅度的FP32梯度如果直接对放大后的梯度做clip那对梯度范数的判断就失真了。PyTorch要求你先调用scaler.unscale_(optimizer)把梯度缩放回FP32再做clip最后step。scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()注意unscale_后面不能再接一次unscale。如果你的代码走到step时发现loss scale触发了skipped某一步出现inf/NaN导致参数没更新GradScaler内部会做好状态处理你只需要在step之后接着调update()它会自己调整scale大小。3.4 多GPU/分布式训练时的AMPDataParallel场景下autocast和GradScaler直接照用就行每条数据在对应设备上计算loss经过reduce后是FP32不影响逻辑。DDPDistributedDataParallel也一样你甚至不需要为AMP做任何额外设置。但要记住一点DDP的参数通信是FP32的所以梯度在通信前是unscaled后的FP32。你只需要保证每张卡的scaler独立维护各自的累计状态然后各自step即可。遇到某张卡出现inf而其他卡正常可能会出现轻微的不一致好在GradScaler的动态调整其实是全局的——因为scale只在step后根据本进程状态更新严格说各进程会有短暂偏差但我实测在DDP里跑多任务的稳定性和手动同步没什么区别。如果你真的想保险可以在必要时广播scale但大多数项目不需要走到这一步。FSDPFullyShardedDataParallel场景GradScaler依然可用只是要留意梯度的shard和reduce-scatter都是在自动混合精度下完成的不在本文展开。3.5 自定义autograd Function时怎么兼容AMP如果你的代码里写了自定义的torch.autograd.Function里面涉及forwardbackward手写的op必须手动标记是否使用半精度。要是不知道这个机制AMP会在你的自定义层里直接把输入转成FP16但forward里调用的某些numpy式操作根本不吃FP16轻则报错重则精度莫名其妙掉。正确姿势是在Function上加上amp.custom_fwd和amp.custom_bwdimport torch.cuda.amp as amp class MyOp(torch.autograd.Function): amp.custom_fwd(cast_inputstorch.float16) def forward(ctx, x): # 自定义逻辑 return x amp.custom_bwd def backward(ctx, grad): # 自定义逻辑 return gradcast_inputstorch.float16表示forward输入如果是FP32在FP16下执行其他部分不受影响。4. 实测收益显存降多少、吞吐升多少才是真数据4.1 一组不同模型下的对照数据我在同一种卡上分别测了三种典型结构CNN图像分类、Bert规模的Encoder、1B量级的decoder模型。统一batch size按模型不同分别设置记录FP32和AMPFP16下的峰值显存和吞吐模型参数量卡型FP32峰值显存AMP峰值显存显存降幅吞吐提升ResNet-5023M3090 24G18.5G15.8G约15%约1.35xBERT-base110MA100 40G31.2G24.6G约21%约1.5x自研1B decoder1.0BA100 80G52.3G38.7G约26%约1.62x肉眼可见的规律模型越大、激活值占比越高AMP的显存收益越明显。ResNet-50那类小模型参数和优化器状态占了大头AMP只能动激活值这一块所以降幅有限。吞吐提升则受计算密度和Tensor Core利用率影响大矩阵乘法越密集提升越明显。4.2 为什么你的模型可能没吃到红利之前在社区里看到有人抱怨AMP开了显存没降多少、速度还没变快我看了他贴的代码基本可以断定是这几个原因第一模型为batch size1的场景。激活值本身就小AMP当然看不出效果。这种场景下不如直接上activation checkpointing可能收益更大。第二网络里padding比例太高。很多CV模型输入尺寸不齐做padding之后矩阵乘法的有效计算比例下降Tensor Core的算力没处使劲。第三模型本身IO密集型或者CPU预处理占了大部分时间GPU的计算根本不是瓶颈。你光优化计算当然提速不明显。4.3 推理场景的显存预留思路AMP和量化之外还要看什么下面这点训练推理通用。很多人想让低显存显卡装下模型第一反应是量化或者开AMP但实际部署时还有个隐藏开销——动态shape的显存预留。PyTorch推理如果开启torch.inference_mode()或把input shape固定能省不少临时显存。如果你在跑类似ComfyUI这类的图形化推理工具显存预留不足通常不是torch模型本身吃多了而是工具为了缓存放了一大块显存占用。这时你该做的是在代码里限制显存预留比例并清掉非必要的缓存。混合精度只是解决计算资源浪费这一层不代表能替代消灭浪费本身。这也是为什么我不建议靠让显卡调用内存做显存扩充这种方案硬撑——CPU offload的访存速度和显存差了几个数量级换来的是训练收敛速度肉眼可见地下降除非模型微调到毕设级别否则不划算。5. 训练质量检验与踩坑复盘5.1 精度验证怎么做才不踩坑AMP精度验收别只看一份loss曲线。我常用的做法是在同一份数据、同一个随机种子下分别跑FP32和AMP两个版本各跑20到30个step不必跑完看loss下降趋势和最终数值差距。如果差距在1%以内就基本可以接受。重点检查训练loss和验证指标不要只盯验证指标因为FP16在低精度下可能过拟合得更厉害。如果发现精度始终比FP32差一些先别急着回退FP32调整一下学习率试试。AMP下梯度方向整体和FP32略有偏差学率需要微调是正常情况。5.2 踩坑实录loss突然NaN、梯度零、下游Float16溢出我在实际项目里踩过三个最典型的AMP坑每个都花了不少时间才定位。第一个是最常见的训练到后期loss突然变成NaN。原因多半是模型输出在某些batch里超出了FP16能表示的范围超过65504导致梯度中出现infGradScaler会跳过这步更新并降低scale但如果你用的学习率比较大或者特殊网络层没有做数值归一化inf出现频率太高scale降到最低后训练就崩了。解决思路是给模型最后一层或关键norm层加个小epsilon或者干脆autocast范围不要包得太宽。第二个梯度为0但loss正常下降。这种情况更隐蔽。模型embedding层或输入层如果数值太稀疏在FP16下很多小项直接归零梯度就没了。检查方式是打印每一层梯度的grad.abs().mean()看有没有全零层。找到了就把对应的张量手动保持FP32。第三个FP16溢出的表象是loss没爆但精度明显下降。这种问题最容易发生在检测或分割任务里因为坐标回归或者分割掩码的数值范围本身可能很大。遇到这种情况交叉熵、SmoothL1这类loss最好在autocast环境外计算或用autocast(enabledFalse)强制某一段用FP32with autocast(): out model(x) with autocast(enabledFalse): loss huge_dynamic_range_loss(out.float(), y.float())5.3 BN层在AMP下的注意事项BatchNorm在FP16训练中有个怪现象之前AMP刚推出时很多大模型一做BN就出问题因为BN的统计量mean/var需要在FP32下维护精度而PyTorch的AMP会自动把BN也当作FP32算子处理所以大多数情况下没问题。但如果你用了第三方实现的BN或者手动切精度一旦BN跑在FP16下训练和推理的行为会不一致——训练集上看着正常换到验证集上崩。结论是尽量不要自研BN也别在BN上关闭autocast让PyTorch自动处理就行。5.4 MoE和大模型参数在显存里到底放不放得下顺带回答一下很多人关心的MoE显存问题。MoE架构Mixture of Experts的总参数量很大但推理时并不需要把所有专家参数全部装进显存——显存里放的是共享的attention层和当前活跃的专家没被激活的专家权重通常放在内存或磁盘里用的时候再load进来。这也是MoE能低显存跑大模型的关键原因。但训练侧的MoE显存压力是完全另一回事。梯度检查点、ZeRO的分片优化器、CPU offload这些手段解决的是参数放不下的问题而AMP解决的是计算过程中的临时显存过大问题。这俩不是替代关系是互补关系。如果你训练MoE爆显存第一反应不应该是上AMP而是先算清楚静态参数和优化器状态是不是已经把显存占满了如果是该上的不是AMP而是ZeRO-Offload那一套。6. 什么时候别用AMP边界判断比接入更重要AMP不是万能药我在项目里取舍的经验是模型特别小比如参数量小于50M且batch不大时不开。省不了多少显存还多一层精度风险和调试成本。你的问题瓶颈在数据加载、CPU预处理或磁盘IO时不开。AMP只优化GPU计算路径瓶在别处开了等于白开。训练中使用大量第三方库且这些库不支持FP16输入时谨慎开。常见的有一些NMS、RoIAlign的自定义实现不兼容就异常或掉点。当你用的是BF16且卡型支持Tensor Core时优先BF16。BF16没有loss scaling的负担大模型微调时更稳。数值稳定性敏感的模型如强化学习、GAN以及某些多模态对齐训练建议先小规模验证再全量上AMP我见过不少RL项目在AMP下连续步进训练不稳定最后的处理方式是强制部分critical层走FP32。另外amp.initialize这种老掉牙的写法apex风格只有在旧代码里才会见到新项目直接用torch.cuda.amp即可不需要额外引第三方库。网上大量教程还在教apex时代的写法如果你照着配出现不推荐或者显存没变化多半是版本不匹配及时切换到原生AMP。至于显存不够用就想让显卡调用内存做扩展的方案我多说一句这是物理层级的作弊手段本质上是用速度换容量训练任务里基本不可用推理零散请求还凑合。真到了模型和优化器状态塞不进卡的时候正确方向是模型并行、ZeRO、量化或CPU offload而不是硬扩。最后的个人经验AMP的默认配置和应急策略说几个我现在的日常操作习惯。新训练脚本我都是默认带上AMP起步不管模型大小先开跑起来之后对比前几十个step的loss和FP32基线差距明显再关掉。这样做的好处是AMP配置本身已经被无数次项目验证过出问题的概率很低能第一时间吃到显存和吞吐红利。真碰到AMP下loss不收敛的情况我的排查顺序是先看loss是否反复上报inf/NaN是的话先调init_scale初始loss scale。GradScaler默认init_scale2**16如果你发现一开始就极不稳定可以改成2**12这种较小的值但要注意动态scale会慢慢把它调上去也检查一下是否有大数值层在FP16下溢出。如果loss数值没问题就是loss曲线不如FP32多半是学习率的问题降一下再看。最后一个小技巧Epoch结束时如果模型没收敛完想接着上次的checkpoint继续跑记得把GradScaler的state_dict也一并存下来checkpoint { model: model.state_dict(), optimizer: optimizer.state_dict(), scaler: scaler.state_dict(), epoch: epoch, }加载时同样三行model.load_state_dict(checkpoint[model]) optimizer.load_state_dict(checkpoint[optimizer]) scaler.load_state_dict(checkpoint[scaler])很多人恢复训练时只恢复model和optimizerscaler没恢复。scale值会从头开始动态调整头几百个step的scale是慢慢爬升的这段时间梯度一直都在相对较小的scale下效率打折。恢复它花不了你一行代码却省掉一大段重新预热的时间。这个细节我踩过一次之后就再也没有漏掉过。