资讯动态

PyTorch AMP混合精度训练实战:显存减半与训练吞吐提升指南

发布时间:2026/9/29 18:35:54 来源:尧图企业网站定制
先说个我前阵子遇到的事儿一个朋友在做LoRA微调显卡是8G显存模型刚加载完就报CUDA out of memory他第一个反应是把batch size从4降到2结果还是炸最后来找我问有没有什么办法能在不大改代码的前提下把显存压下来。我反手就给他开了AMPbatch size恢复成4还顺手把训练吞吐提了30%以上。这就是PyTorch AMP混合精度训练的典型应用场景在训练和微调阶段用一种很轻量的方式同时解决显存爆炸和算力跑不满两个问题。AMP不是把模型简单粗暴地变成半精度而是让PyTorch在运行过程中自动为合适的算子切换到FP16同时保留关键部分的FP32精度。这篇内容就是把我实际用AMP的经验完整梳理一遍——适合正在用裸PyTorch写训练循环、想搞懂底层机制、或者用Lightning/Accelerate但不确定底层到底怎么跑的朋友。1. 显存焦虑与吞吐瓶颈为什么混合精度成了刚需我自己见过太多“显存不够”的求助帖了但很多人没搞明白显存到底被谁吃了。一个7B模型训练时的显存开销大概是这样的模型权重7B × 4字节 ≈ 28GBFP32梯度同样是28GB优化器状态AdamW需要保存一阶动量m和二阶动量v各一份约56GB激活值前向计算中为了反向传播而保留的中间张量这个数量级随batch size和模型结构浮动经常是训练中最大的开销这还没算上真正的激活值显存。所以8G显存想全量训练7B模型就算用FP32也完全不可能LoRA之所以能跑是因为被训练的额外参数很少优化器状态和梯度开销小了一大块剩下的显存大头就是激活值。再谈吞吐。现代NVIDIA显卡都带Tensor Core专门为FP16这类低精度的矩阵乘法和卷积做了加速。以A100为例FP16的峰值算力大概是FP32的两倍。如果你整个训练过程都在FP32里跑等于让Tensor Core在旁边干瞪眼算力根本没发挥出来。所以AMP解决的就是两个核心问题显存占用和计算吞吐。它不是把训练代码推倒重写而是从框架层面自动做精度调度把能安全用FP16的部分交给Tensor Core去跑把敏感的数值计算留在FP32。这也是为什么AMP可以成为训练优化的默认选项——改动小、收益大、风险相对可控。2. FP16的精度陷阱溢出、下溢与损失缩放机制拆解说到AMP就绕不开FP16。但FP16并不是一个“更好”的精度它有很大的数值表示盲区。FP16用16位二进制表示一个数其中1位符号、5位指数、10位尾数。它的最大值是65504最小正规格化数大约是6e-5。FP32呢指数有8位范围可以到1e-38到3.4e38。看数字可能不够直观我打个比方FP32能表示的数值范围相当于地球到太阳的距离而FP16能表示的范围可能就相当于一张桌子那么长。这不是夸张是几个数量级的差距。那问题来了神经网络训练里的梯度数值通常在1e-3到1e-8之间波动。对FP32来说1e-8这种数量级完全没问题但对FP16低于6e-5的梯度直接就会下溢成0。一个梯度为0的参数就等于这次更新什么都没做。所以“全模型切成FP16”这条路基本走不通。早在多年前的混合精度实践里业界就总结出了三件套FP16存储与计算中间量大幅减少显存并利用Tensor Core。FP32主权重master weight模型参数的“正式版本”依然保存在FP32里每次更新都在FP32上完成然后再把FP32副本转成FP16供前向和反向使用。这样可以避免反复累加时的小数误差累积。损失缩放Loss Scaling这是最巧妙的一步。既然梯度太小会下溢成0那就在反向传播之前先给loss乘一个大数比如65536。因为链式法则梯度在反向传播过程中也会同等地被放大这样原本在FP16下会变0的梯度就能被正常表示了。等梯度计算完真正更新参数之前再把它除回去。关键点在于为什么缩放因子偏偏喜欢用2的幂因为FP16本身的存储结构就是指数位尾数位乘以2的幂只是在指数部分做加法尾数一个位都不会丢衰减回去的时候也是完全精确的不会引入额外误差。PyTorch的GradScaler实现的是动态损失缩放初始scale通常是2的16次方训练过程中一旦检测到梯度出现inf或nan就判定数值溢出马上把scale减半回退如果连续很多步都正常就逐渐把scale翻倍继续尝试。整个过程是自动的你只需要在训练循环里按要求调用API。3. 新旧API与正确姿势torch.amp如何取代torch.cuda.ampAMP在PyTorch里有两个核心APIautocast自动类型转换上下文和GradScaler梯度缩放器。先说版本演进。torch.cuda.amp是PyTorch 1.6时代加入的API设计得很直白但名字里带了cuda天然跟CUDA绑死。PyTorch 2.x开始推出统一的torch.amp支持CUDA、CPU等不同设备类型。所以新代码我建议直接用torch.amp写法如下from torch.amp import autocast, GradScaler scaler GradScaler(cuda, init_scale2**16, growth_factor2.0, backoff_factor0.5, growth_interval2000)torch.amp.autocast接收第一个参数是设备类型。以前用torch.cuda.amp.autocast()现在建议写成torch.amp.autocast(cuda, dtypetorch.float16)。autocast的工作机制可以理解为一个“按需调度器”。在它的上下文范围内PyTorch遇到不同的算子会自动选择输入输出精度而不是把所有的tensor都变成FP16。这类算子会被自动调度到FP16矩阵乘法matmul、bmm、addmm等线性层Linear卷积层Conv1d/2d/3d嵌入层EmbeddingLSTM等常见循环层但另一些算子会强制留在FP32最典型的是LayerNorm、BatchNorm、Softmax、CrossEntropyLoss、Exp这类对数值稳定性要求高的操作。为什么会这样因为这些操作对精度极其敏感一旦输入被放到FP16误差会非线性放大尤其在Transformer类模型里LayerNorm跑在FP32几乎是最低要求。这里要顺手纠正一个初学AMP最容易犯的错不要手动把输入tensor转成half半精度。很多人觉得“既然要混合精度那我先把数据转成FP16”结果精度掉得一塌糊涂。autocast会自动处理输入tensor的精度转换你手动转了反而可能让某些算子接收到错误的dtype甚至直接报类型不匹配。在autocast上下文里你该写的代码和平时完全一样模型接收FP32输入然后PyTorch自动调度。4. 手把手把标准训练循环改成AMP最小可复现改造模板如果训练循环是你手写的改造只要加5行左右的代码。下面是一段标准的PyTorch训练循环改成AMP后的样子import torch from torch.amp import autocast, GradScaler model MyModel().cuda() optimizer torch.optim.AdamW(model.parameters(), lr1e-4) criterion torch.nn.CrossEntropyLoss() scaler GradScaler(cuda, init_scale2**16) for batch in dataloader: optimizer.zero_grad() # 前向传播和loss计算放进autocast上下文 with autocast(cuda, dtypetorch.float16): outputs model(batch[input_ids], batch[attention_mask]) loss criterion(outputs, batch[labels]) # 反向传播用scaler.scale(loss)替代loss.backward() scaler.scale(loss).backward() # 必须先unscale_再clip梯度 scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()逐行拆解这几个关键变化第一为什么loss和后向计算之间要用scaler.scale(loss).backward()因为我们要给loss乘以缩放因子让梯度在FP16可表示范围内。这里有个容易忽略的点loss.backward()是在放大后的loss上进行的所以梯度本身也被放大了。scaler.step(optimizer)内部会判断梯度是否有效有没有inf/nan如果有效就先做一次unscale再真正执行优化器更新。第二为什么unscale_必须在clip_grad_norm_之前梯度裁剪gradient clipping的阈值是基于真实梯度数值的。如果不先把缩放因子除掉你裁剪的其实是被放大了65536倍的梯度那裁剪几乎等于没有裁甚至会让梯度幅度判断完全失真。我在刚接触AMP时就踩过这个坑损失曲线稳定不掉点但验证集指标特别差后来发现是裁剪计算全错了。正确顺序就是上面代码里的unscale_→clip_grad_norm_→step。第三为什么scaler.update()放在最后它负责根据这一轮是否出现inf/nan来动态更新缩放因子。如果梯度正常它可能会把scale值往上抬如果有溢出现象就往下压。它必须在这一轮优化器更新完成之后执行因为它是为下一轮训练做准备的。还有一个细节如果你想在日志里打印真实loss不能直接用loss.item()因为现在loss是经过scale的。需要除回当前缩放因子current_scale scaler.get_scale() real_loss loss.item() / current_scale如果你用的是torch.compile或Lightning这类高层框架一般会有对应的内置AMP开关但底层逻辑和这套完全一样。理解裸PyTorch的AMP改造方式能帮你在框架封装不透明的时候快速定位问题。5. 精度掉点现场排查黄金对照实验与调参思路AMP最大的心理负担就是“会不会掉点”。我在用AMP的过程中遇到过掉点但绝大多数情况都不是AMP本身的问题而是其他环节被精度变化放大了。这里分享一个我经常用的排查流程。第一步先做黄金对照实验在完全相同的随机种子、相同batch数据、相同学习率下分别用纯FP32和AMP各训练100步记录loss曲线和验证指标。AMP开启后loss曲线出现稍微波动是正常的毕竟计算顺序和数值路径都不一样但如果100步后两个loss有明显分歧比如FP32已经稳定在2.0、AMP还在3.5徘徊那就需要继续排查。排查顺序我一般是这样检查是不是手动转换了输入dtype。autocast上下文之外把输入转成了half大概率会导致问题。检查自定义loss函数。如果你写了一个很复杂的自定义loss里面涉及不支持的算子有时候会自动走下推逻辑、有时会报错。官方支持列表之外的操作最好先确认一下在autocast下的行为。最简单的验证方法是把自定义loss里的每个算子都单独拎出来分别用FP32输入和FP16输入跑一下看输出结果是否一致。观察scaler.get_scale()的走势。如果缩放因子一直在回退说明模型训练中频繁出现inf/nan这通常不是AMP的问题而是模型本身或初始化数值就不稳定。可以尝试降低学习率或者调整初始化方式。学习率要不要动。很多人在开AMP的同时顺手改大了batch size却发现掉点了。这不是AMP的锅而是batch size增大后学习率没做相应缩放。一般建议batch size翻倍学习率要么线性翻倍要么按平方根比例缩放先做小规模实验确认。如果FP16怎么调都救不回来精度换BF16试试。BF16的全称是bfloat16指数位和FP32一样都是8位数值范围和FP32几乎一样因此不会出现FP16那种剧烈的下溢问题代价是尾数精度低。对很多大模型训练来说BF16的稳定性比FP16好很多尤其适合已经对FP16不友好的Transformer结构。在autocast里只需要把dtype改成torch.bfloat16即可而且BF16不需要GradScaler——因为它的动态范围和FP32接近很难下溢。细节是BF16在30系、40系等较新的NVIDIA显卡上支持良好更老的卡需要确认算力是否匹配。还有一个非常重要的经验不要拿AMP后的模型精度和原模型做一次“决赛”式的对比因为一次训练的波动本身就很大。多做几组不同seed的小实验再下结论更靠谱。6. AMP救不了的那部分显存激活值、优化器状态与BatchNorm很多朋友开了AMP以后发现显存确实降了但降到一定程度就不动了于是跑过来问要不要把模型也.half()一下。这里得说清楚一个容易混淆的点AMP降显存的主要来源是激活值activation和反向传播的中间张量而不是模型权重。举个例子一个Transformer层里线性层的输入和输出都会成为反向传播需要的中间张量。在FP32下batch size 16、序列长度2048、隐藏层4096光是一个tensor就是16×2048×4096×4字节约512MB。Tensor转成FP16后直接减半变成256MB。一层省一点几十层堆下来就是好几个GB。但模型权重呢AMP默认还是用FP32保存的“正式版本”因为每次更新都在FP32上进行然后临时转成FP16参与前向计算。所以权重本身的内存并不会因为开启AMP而自动减半。如果你真正的大头在优化器状态——比如你全量训练一个大模型AdamW的m和v加起来就是两倍参数量——那AMP就是“杯水车薪”。这也是为什么LoRA AMP这个组合非常常见LoRA把可训练参数压到极小优化器状态也就压到极小显存大头变成激活值此时AMP就能发挥最大作用。我那个朋友的8G显存LoRA场景不开AMP时batch size 2都抖开AMP后batch size 4稳得很。另一个细节是BatchNorm。很多人以为开AMP之后所有层都会变成FP16但实际情况是BatchNorm在autocast下依然会用FP32计算——这是PyTorch的默认策略。为什么因为BatchNorm的本质是在一个batch的维度上做归一化它需要计算均值和方差这些统计量对数值精度特别敏感一旦进了FP16统计结果就不太稳了训练和推理的差异也会被放大。LayerNorm同理。所以如果你的模型是纯卷积网络层里没有太多BatchNormAMP的显存收益可能非常可观但如果是Transformer里面有大量LayerNorm和Softmax这些部分会保持在FP32实际显存降幅就没那么夸张。显存还是不够怎么办我建议叠加组合拳先开AMP再开gradient checkpointing用时间换显存把少数前向中间张量不存储、反向时重算再配合LoRA/低秩分解最后才考虑8bit优化器或者4bit量化微调。这一步一步做下来8G显存跑一些中等级别的模型是完全可行的。判断AMP到底帮你降了多少显存别凭感觉用代码说话torch.cuda.reset_peak_memory_stats() # 你的训练循环跑几十步 peak_memory torch.cuda.max_memory_allocated() print(f峰值显存: {peak_memory / 1024**3:.2f} GB)同样的位置分别在FP32和AMP下各跑一次就能得到实际降幅。我第一次测自己项目的时候峰值显存从11.2GB降到7.8GB那感觉是真的舒爽。7. 实测收益与适用边界显存降幅、吞吐增幅与不要神化AMP说到实测收益我得先泼一盆冷水AMP不是你加上去就一定快30%的银弹。它的实际收益高度依赖你的模型结构、batch size、显卡算力以及显存瓶颈在哪。我把常见的几类场景做了个经验总结注意这是经验区间不是普适承诺场景显存降幅吞吐增幅备注中大规模Transformer全量微调20%~40%40%~80%激活值占比高Tensor Core利用率提升明显LoRA/AdaLoRA微调30%~50%20%~50%优化器状态变小AMP作用被放大卷积网络图片分类20%~30%30%~60%BatchNorm保持FP32收益略低于Transformer小模型/单层MLP10%以下10%以下计算密度太低通信和调度开销占比高CPU训练无无CPU上AMP基本没有加速效果少折腾在A100这类算力强、显存带宽高的卡上AMP的收益会被放大在老消费级卡上如果你的batch size已经很小比如batch size 1那AMP带来的可能主要是显存降低吞吐提升不一定明显。我建议所有正在纠结“要不要开AMP”的朋友都做一个十分钟小实验固定好数据、模型、随机种子用FP32和AMP各跑30个step记录两个指标torch.cuda.max_memory_allocated()峰值显存每秒钟处理的样本数这个小实验能帮你判断AMP对你当前场景的实际收益。如果显存降了30%、吞吐提了40%那几乎没有什么理由不开如果收益不到10%那可以把精力花在数据加载、batch size调整或模型结构优化上。还有一个容易忽略的好处显存降了你就有空间把batch size调大。batch size越大GPU利用率越高吞吐还能再上一个台阶。但注意batch size变大后梯度方向会更稳定学习率可能需要相应调整这又回到了上一节说的调参问题。8. 踩坑记录DDP、梯度裁剪、动态缩放与dtype一致性最后这部分是我觉得最有价值的因为光看文档你不会知道这些坑有多深。我按实战中遇到过的坑挨个说。坑1梯度裁剪必须放在unscale_之后。这个我在前面已经强调过但值得再重复一次。PyTorch官方文档也明确写了如果你用了GradScaler必须先调用scaler.unscale_(optimizer)再执行梯度裁剪。否则裁剪的是缩放后的梯度等于白裁。有段时间我把scaler.step(optimizer)和scaler.update()写对了但裁剪顺序写反了结果训练loss曲线看起来很正常就是验证集指标不动排查了我大半天。坑2DDP下每个rank都要独立创建scaler且step不能跳过。使用DistributedDataParallel时每个进程有自己的模型副本也需要有自己的GradScaler实例。更关键的是scaler.step()这个调用必须在所有rank上同步执行——因为DDP在反向传播后会做梯度同步如果一个rank因为梯度无效而跳过了step其他rank的通信状态会错乱轻则训练不稳定重则卡死。写DDP代码时别自作聪明地在某个rank上跳过scaler.step()。如果某个rank检测到inf/nan让scaler.step()自己处理它会跳过优化器更新并更新缩放因子。坑3梯度累积时要小心缩放因子。假设你设置accumulation_steps4也就是4个小batch的梯度累加后再更新一次。AMP下每个小batch的loss都会被scaler.scale()这些缩放的梯度会累加到参数梯度上。这本身没有问题但有个微妙之处如果其中一个小batch的梯度溢出成了inf整个累加结果都废了。遇到这种情况最常见的手段是把无效的那个mini-batch从计算图中分离重新forward一次复杂而且烦人。我的建议是梯度累积场景下优先用BF16因为它几乎没有溢出风险如果必须用FP16就把accumulation step的batch里最后一个step放在scaler.step()之前特别留意。坑4自定义的CUDA kernel或flash-attn类库可能不认识autocast。有些第三方高效算子库在autocast上下文内并不一定乖乖处理FP16输入可能直接报dtype不匹配或者走了一个没优化的fallback。排查方法是单步debug把出问题的算子的输入输出dtype都打出来。如果发现某个自定义算子在FP32下正常、FP16下异常可以先在autocast外面计算这个部分或者手动用torch.cuda.amp.autocast(enabledFalse)局部关闭精度调度。坑5loss本身是FP16时不要再用GradScaler。虽然这种情况比较少见但如果你写了自定义loss返回的是FP16 tensorscaler.scale(loss)会直接报错因为缩放器要求loss是FP32或FP64。我自己遇到过类似问题解决方法是把loss统一在后面转成FP32或者干脆保证loss的计算回到FP32再交给scaler。坑6eval阶段不是必须开GradScaler但建议开autocast。推理/验证阶段没有反向传播就不用GradScaler但可以继续用autocast来降低验证阶段的显存占用有些算子在FP16下还更快一些。注意一点如果模型里含BatchNorm训练和推理的统计量使用方式不同推理时用autocast也尽量确认输出精度符合预期。最后说一句我的实际体会AMP在现在的PyTorch训练流程里已经是一个“性价比极高”的默认选项了。它不解决所有问题优化器状态太大、激活值过大的结构性问题它都救不了但作为第一板斧去砍显存和提吞吐几乎稳赚不赔。如果你还从来没试过找一个训练循环加上那五行代码跑30个step对比一下数据你自己就知道答案了。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑