拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

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

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

先说个我前阵子遇到的事儿:一个朋友在做LoRA微调,显卡是8G显存,模型刚加载完就报CUDA out of memory,他第一个反应是把batch size从4降到2,结果还是炸,最后来找我问有没有什么办法能在不大改代码的前提下把显存压下来。我反手就给他开了AMP,batch size恢复成4,还顺手把训练吞吐提了30%以上。

这就是PyTorch AMP混合精度训练的典型应用场景:在训练和微调阶段,用一种很轻量的方式同时解决显存爆炸和算力跑不满两个问题。AMP不是把模型简单粗暴地变成半精度,而是让PyTorch在运行过程中自动为合适的算子切换到FP16,同时保留关键部分的FP32精度。这篇内容就是把我实际用AMP的经验完整梳理一遍——适合正在用裸PyTorch写训练循环、想搞懂底层机制、或者用Lightning/Accelerate但不确定底层到底怎么跑的朋友。

1. 显存焦虑与吞吐瓶颈:为什么混合精度成了刚需

我自己见过太多“显存不够”的求助帖了,但很多人没搞明白,显存到底被谁吃了。一个7B模型,训练时的显存开销大概是这样的:

  • 模型权重:7B × 4字节 ≈ 28GB(FP32)
  • 梯度:同样是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”这条路基本走不通。早在多年前的混合精度实践里,业界就总结出了三件套:

  1. FP16存储与计算中间量,大幅减少显存并利用Tensor Core。
  2. FP32主权重(master weight):模型参数的“正式版本”依然保存在FP32里,每次更新都在FP32上完成,然后再把FP32副本转成FP16供前向和反向使用。这样可以避免反复累加时的小数误差累积。
  3. 损失缩放(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.amp

AMP在PyTorch里有两个核心API:autocast(自动类型转换上下文)和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_scale=2**16, growth_factor=2.0, backoff_factor=0.5, growth_interval=2000)

torch.amp.autocast接收第一个参数是设备类型。以前用torch.cuda.amp.autocast(),现在建议写成torch.amp.autocast('cuda', dtype=torch.float16)。

autocast的工作机制,可以理解为一个“按需调度器”。在它的上下文范围内,PyTorch遇到不同的算子会自动选择输入输出精度,而不是把所有的tensor都变成FP16。这类算子会被自动调度到FP16:

  • 矩阵乘法(matmul、bmm、addmm等)
  • 线性层(Linear)
  • 卷积层(Conv1d/2d/3d)
  • 嵌入层(Embedding)
  • LSTM等常见循环层

但另一些算子会强制留在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(), lr=1e-4) criterion = torch.nn.CrossEntropyLoss() scaler = GradScaler('cuda', init_scale=2**16) for batch in dataloader: optimizer.zero_grad() # 前向传播和loss计算放进autocast上下文 with autocast('cuda', dtype=torch.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_norm=1.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徘徊,那就需要继续排查。

排查顺序我一般是这样:

  1. 检查是不是手动转换了输入dtype。autocast上下文之外把输入转成了half,大概率会导致问题。
  2. 检查自定义loss函数。如果你写了一个很复杂的自定义loss,里面涉及不支持的算子,有时候会自动走下推逻辑、有时会报错。官方支持列表之外的操作,最好先确认一下在autocast下的行为。最简单的验证方法是:把自定义loss里的每个算子都单独拎出来,分别用FP32输入和FP16输入跑一下,看输出结果是否一致。
  3. 观察scaler.get_scale()的走势。如果缩放因子一直在回退,说明模型训练中频繁出现inf/nan,这通常不是AMP的问题,而是模型本身或初始化数值就不稳定。可以尝试降低学习率,或者调整初始化方式。
  4. 学习率要不要动。很多人在开AMP的同时顺手改大了batch size,却发现掉点了。这不是AMP的锅,而是batch size增大后学习率没做相应缩放。一般建议:batch size翻倍,学习率要么线性翻倍,要么按平方根比例缩放,先做小规模实验确认。
  5. 如果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同理。

所以如果你的模型是纯卷积网络,层里没有太多BatchNorm,AMP的显存收益可能非常可观;但如果是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曲线看起来很正常,就是验证集指标不动,排查了我大半天。

坑2:DDP下每个rank都要独立创建scaler,且step不能跳过。使用DistributedDataParallel时,每个进程有自己的模型副本,也需要有自己的GradScaler实例。更关键的是,scaler.step()这个调用必须在所有rank上同步执行——因为DDP在反向传播后会做梯度同步,如果一个rank因为梯度无效而跳过了step,其他rank的通信状态会错乱,轻则训练不稳定,重则卡死。写DDP代码时,别自作聪明地在某个rank上跳过scaler.step()。如果某个rank检测到inf/nan,让scaler.step()自己处理,它会跳过优化器更新并更新缩放因子。

坑3:梯度累积时要小心缩放因子。假设你设置accumulation_steps=4,也就是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(enabled=False)局部关闭精度调度。

坑5:loss本身是FP16时,不要再用GradScaler。虽然这种情况比较少见,但如果你写了自定义loss返回的是FP16 tensor,scaler.scale(loss)会直接报错,因为缩放器要求loss是FP32或FP64。我自己遇到过类似问题,解决方法是把loss统一在后面转成FP32,或者干脆保证loss的计算回到FP32再交给scaler。

坑6:eval阶段不是必须开GradScaler,但建议开autocast。推理/验证阶段没有反向传播,就不用GradScaler,但可以继续用autocast来降低验证阶段的显存占用,有些算子在FP16下还更快一些。注意一点:如果模型里含BatchNorm,训练和推理的统计量使用方式不同,推理时用autocast也尽量确认输出精度符合预期。

最后说一句我的实际体会:AMP在现在的PyTorch训练流程里,已经是一个“性价比极高”的默认选项了。它不解决所有问题,优化器状态太大、激活值过大的结构性问题它都救不了,但作为第一板斧去砍显存和提吞吐,几乎稳赚不赔。如果你还从来没试过,找一个训练循环,加上那五行代码,跑30个step对比一下数据,你自己就知道答案了。

返回列表