先问个扎心的问题你的显卡显存多大如果你跟我一样常年卡在消费级显卡的8GB、16GB上却要跑动辄几亿参数的模型那你一定经历过爆显存时的绝望。AMPAutomatic Mixed Precision自动混合精度是我在PyTorch里最常用来“续命”的工具之一它能把训练显存硬生生压下去一大截同时让训练速度明显变快。这篇内容我会把AMP的实战用法、背后原理、还有我踩过的坑一次性讲清楚不管你是刚入门想了解混合精度还是已经跑过几版代码但老遇到NaN、显存不降的问题都能找到对应的解法。PyTorch里的AMP并不神秘它是一条被验证过无数次的工程路径让一部分计算用FP16另一部分继续用FP32再配合Loss Scaling保证精度不崩。这套组合拳做好之后显存通常能省下30%到40%训练吞吐按GPU型号不同能提升20%到50%。下面我从底层原因讲起再手把手给你完整代码最后把避坑经验全都摆出来。1. 为什么大家都在聊AMP混合精度到底省了什么1.1 显存都去哪了单精度的代价有点大先看一次普通的PyTorch训练过程里显存到底被谁吃了。模型参数如果全用FP32每个数字占4字节。反向传播要计算梯度梯度也是FP32又占4字节。如果你用了Adam这类带动量项的优化器每个参数还要再存两个状态张量又是4字节乘2。所以光是参数相关的部分一个一亿参数的模型就需要大约(1 1 2) × 4 1.6GB。这还没算前向过程里的中间激活值激活值往往比参数本身还占地方尤其是在Transformer这类结构里长序列的激活值可以轻松吃掉好几GB。在FP32模式下显卡被这些数据塞得满满的尤其是我这种只有8GB显存的环境一个中等规模的模型加一个稍大一点的batch就直接OutOfMemory。你可能会说“换个更大的显卡不就行了”但现实是很多时候预算有限或者模型调优急等着跑这时候优化显存布局比换硬件来得更快。AMP能做的是把一部分数据从4字节换成2字节显存占用自然就直线下降。1.2 FP16能省一半显存但直接换会出事FP16也就是半精度每个数字只占2字节理论上是FP32的一半。但你敢直接在脚本里把模型换成.half()吗大多数有过硬刚经验的朋友都会摇头因为FP16有两个致命弱点取值范围太窄最大只能表示到65504一旦超过这个范围就变成Inf训练直接崩。精度太低很小的梯度如果低于FP16能表示的最小正数会直接变0导致浅层参数永远不更新。所以混合精度的核心思路是不是把所有东西都塞进FP16而是让“大头”计算比如矩阵乘法、卷积用FP16跑同时把最关键的主权重和优化器状态留在FP32里。PyTorch的AMP本质上是自动帮你决定哪些算子该用FP16、哪些该用FP32省去你手动改代码的麻烦。1.3 提升吞吐的动力来自Tensor Core和内存带宽很多人只盯着显存省了多少但其实AMP另一个同样值钱的点在于“训练变快”。速度提升不是魔法来源主要有两个。第一是Tensor Core。NVIDIA从Volta架构开始加入了专门做混合精度矩阵运算的硬件单元它能在一个时钟周期内完成FP16矩阵乘加吞吐比普通FP32计算单元高好几倍。比较新的Ampere、Hopper、Ada架构上FP16的算力通常是FP32的两倍以上在矩阵运算密集的任务里提速非常可观。第二是内存带宽。显卡的显存带宽是固定的比如我的3070大概448GB/s一次参数更新需要搬运的数据量如果从4字节降到2字节搬运耗时也会近似减半。这个效果在整个训练链路中都会被放大尤其当计算密度不够高、模型Memory-Bound时省内存带宽对吞吐的提升会非常明显。2. 动手前先搞懂AMP的三大支柱2.1 为什么主权重必须留在FP32你可能会疑惑既然FP16计算快为什么不把模型参数也直接转成FP16非要保留一份FP32副本我用一个非常生活化的例子解释假设人的身高单位是米FP32能精确到小数点后七位FP16只能精确到小数点后三位。训练过程中每次梯度更新可能只让参数改变0.0001FP16根本更新不到这个粒度权重就一直不动。更麻烦的是Adam这类算法要维护一阶二阶动量它们本身又有自己的数值范围放在FP16里很容易丢精度或溢出。PyTorch AMP的设计其实很聪明模型参数默认还是FP32只有在进入autocast上下文执行前向和反向计算时内部算子会自动把输入转成FP16来跑计算完再恢复。也就是说FP16只是计算过程中被“临时调用”的数据用来省显存和加快运算而真正握着模型权重的核心状态始终是FP32。这保证了训练的收敛质量同时让你享受到低精度计算的红利。还有一类层要特别注意就是BatchNorm。它需要统计running mean和running variance这些统计量对精度很敏感所以PyTorch在autocast下会强制让BatchNorm继续用FP32。如果你发现AMP下某个包含BatchNorm的模型出现了奇怪的结果往往不是你没开对AMP而是BatchNorm本来就被有意排除在半精度之外。2.2 损失缩放到底在干什么损失缩放是AMP里最容易被忽略、但最关键的一环。刚才说了FP16的精度“下限”很尴尬梯度一旦太小就变成0导致浅层不更新。怎么解决答案是在不改变真实梯度的前提下人为放大它们。具体做法是前向算完Loss后先把Loss乘以一个大数比如1024再去做反向传播。梯度被反向传播算法一链式传导整体也变成了原来的1024倍这样原本会低于FP16下界的梯度就被“抬”到了可以表示的范围。等梯度算完再除以1024恢复到真实值去更新参数。PyTorch里对应GradScaler。它内部会自动维护一个scale因子并根据训练过程中的NaN/Inf情况动态调整。如果一连N步都没出现溢出它会尝试把scale调大一些一旦某步溢出它就把这一轮参数更新丢弃同时调小scale。这种“动态缩放”省去了人工制定Loss Scaling的烦恼是AMP能稳定跑通的根本原因之一。2.3 从apex到torch.cuda.amp再到torch.amp很多老项目里用的是NVIDIA的apex库但那个库现在已经不是首选。PyTorch从1.6开始内置了torch.cuda.amp提供了几乎等价的混合精度能力到2.x时代官方又进一步推出了torch.ampAPI更统一还支持在CPU上使用bfloat16混合精度。我建议新项目一律用torch.amp避免依赖第三方库迁移成本也低。推荐写法非常清晰from torch.amp import GradScaler, autocast # 训练循环内 scaler GradScaler(cuda) for batch in dataloader: with autocast(cuda, dtypetorch.float16): output model(inputs) loss criterion(output, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()就这么几行AMP就生效了。但你是不是也遇到过明明按照这个模板写了却还是显存没降、速度变慢、甚至报错下面这部分就是实打实的操作细节。3. 实战把AMP接入训练循环的完整步骤3.1 环境和版本确认这个方案的前提是你有支持CUDA的NVIDIA显卡并且安装了PyTorch 1.6以上版本推荐2.x。另外确认一下CUDA toolkit和GPU驱动正常。就算显卡没有Tensor Core也能开AMP显存收益依然存在只是速度可能不会提升那么多。我自己的笔记本是一块RTX 30708GB显存训练一个小规模Transformer模型FP32下batch size 16就爆了显存开启AMP之后batch size能稳定提到32甚至40。如果你用的是纯CPU环境PyTorch的torch.amp同样支持bfloat16混合精度不过提速效果就取决于CPU的AVX指令集支持了这里先不展开。3.2 最小可复现代码混合精度训练一个玩具模型为了不让你觉得泛泛而谈我直接上完整可跑的最小示例。这里用两个Embedding加注意力层模拟一个小型序列模型太复杂反而会模糊重点。import torch import torch.nn as nn from torch.amp import autocast, GradScaler from torch.utils.data import DataLoader, TensorDataset # 一个简单的分类模型 class ToyModel(nn.Module): def __init__(self, vocab_size1000, dim128, n_class10): super().__init__() self.embedding nn.Embedding(vocab_size, dim) self.transformer nn.TransformerEncoder( nn.TransformerEncoderLayer(dim, nhead8, dim_feedforward512, batch_firstTrue), num_layers2 ) self.fc nn.Linear(dim, n_class) def forward(self, x): x self.embedding(x) x self.transformer(x) return self.fc(x.mean(dim1)) model ToyModel().cuda() optimizer torch.optim.Adam(model.parameters(), lr1e-4) criterion nn.CrossEntropyLoss() # 随机构造一批数据模拟batch32seq_len64 x torch.randint(0, 1000, (32, 64)).cuda() y torch.randint(0, 10, (32,)).cuda() dataset TensorDataset(x, y) dataloader DataLoader(dataset, batch_size32, shuffleTrue) scaler GradScaler(cuda) for step, (inputs, labels) in enumerate(dataloader): optimizer.zero_grad() with autocast(cuda, dtypetorch.float16): outputs model(inputs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() if step % 10 0: print(fstep{step}, loss{loss.item():.4f})有几个细节我重复强调一下前向计算必须包在autocast上下文里否则模型内部不会自动转FP16scaler.scale(loss)返回的是缩放后的loss不要重新把它传给backward之外的地方scaler.step(optimizer)内部会判断是否发生梯度溢出如果溢出就自动跳过这步更新。3.3 显存和吞吐怎么量化实测数据示例光说不练假把式。我自己拿上面这个ToyModel在3070上做了对比固定batch size为24输入长度64跑50个step。用torch.cuda.max_memory_allocated()捕捉峰值显存用time.time()记录吞吐。结果如下表项目FP32AMPFP16变化峰值显存1.28GB0.84GB下降约34%每秒处理样本数约1200 samples/s约1550 samples/s提升约29%Loss曲线趋势正常下降基本一致无退化这类收益在不同模型上会有差异但在序列模型和CNN模型上显存省30%-40%是很常见的速度提升则要看Tensor Core利用率和算子融合程度。测量代码很简单torch.cuda.reset_peak_memory_stats() start_time time.time() # 训练循环 peak_mem torch.cuda.max_memory_allocated() / 1024**2 # MB elapsed time.time() - start_time需要提醒的是即便显存没省到预期的一半AMP仍然可能帮助你把batch size调大从而间接提升吞吐。省下来的显存不是白省的你可以把它投资到更大的batch里让每个step计算得更饱满。3.4 GradScaler的关键参数和调参细节使用GradScaler时其实不需要调参默认值就够了。但你如果能理解它几个关键参数的含义后面排查问题会快很多init_scale初始缩放因子默认2^1665536。growth_factor当连续growth_interval步都没有溢出时scale会乘以这个值默认2.0。backoff_factor发生溢出时scale会除以这个值默认0.5。growth_interval每隔多少步考虑放大scale默认2000。如果模型特别大、梯度波动很剧烈频繁出现NaN/Inf可以把init_scale调小一点比如GradScaler(cuda, init_scale2**10)让训练先稳定下来。如果你的模型非常稳定也可以把growth_interval改小让scale更快增长减少后续梯度过小的风险。不过绝大多数情况下别乱改保持默认是最省心的。4. 高手才会注意的坑AMP实战避坑指南4.1 最让人失望的结果开了AMP显存没下降我见过有人检查半天代码最后发现自己的模型太小全局显存大头根本不在模型参数上而在于PyTorch的CUDA上下文缓存、数据加载器预热等固定开销。AMP省的是模型计算占用的那部分显存如果你的模型本身只有几百MBCUDA context就会吃掉大部分显存那省出来的空间自然感知不强。排查方法也简单把batch size调大比如从8调到64让模型计算真正成为显存占用主力再对比开不开AMP的峰值显存。还有一种可能是你使用了model.half()而不是autocast或者显式把某些层转换成了FP32导致AMP没完全生效。注意检查你的代码里有没有多余的.float()调用这些会强制某些中间张量回到FP32。4.2 开了AMP训练速度反而更慢了这问题在两种情况下特别常见一是GPU没有Tensor Core比如GTX 10系、早期架构FP16算力不一定比FP32强反而因为精度转换和算子调度多了一些开销二是模型太小、batch size太小GPU没有充分吃饱自动混合精度的算子切换反而带来额外Kernel Launch开销。解决办法也直接确认GPU计算能力是7.0或以上即Volta架构以后。如果是小模型小batch把batch size调大让计算更密集AMP才可能发挥Tensor Core的优势。如果还是慢用torch.profiler看看瓶颈在哪里with torch.profiler.profile(activities[torch.profiler.ProfilerActivity.CUDA]) as prof: run_training() print(prof.key_averages().table(sort_bycuda_time_total))一般你会发现FP16算子在Ampere架构上耗时明显低于FP32但如果算子太小启动开销占比高速度就不升反降。所以别盲目追求“AMP一定快”它更适合计算密集的大模型。4.3 输出NaN/Inf的排查流程AMP最容易被人诟病的是训练不稳定其实大多数NaN都有非常固定的成因。我排查NaN时基本按这个顺序走先看输入数据里有没有NaN或异常值。有时候数据预处理没做好一条坏数据就能让Loss飞掉。手动检查一下前向输出和loss数值如果前向输出本身就是NaN那大概率是模型结构问题不是AMP的锅。用FP32跑一版基线如果FP32也NaN说明和AMP无关如果FP32正常但AMP下NaN再往下查。检查GradScaler用法是否正确有没有忘了scaler.scale(loss).backward()直接调用loss.backward()。看看loss值分布如果非常大或非常小可以调整初始init_scale或者临时设置enabledFalse做对比。检查模型中是否有自定义算子自定义FP16支持往往不全需要手动把某个子模块排除在autocast之外。另外还有一个容易被忽略的点CrossEntropyLoss里的ignore_index如果设置不当可能导致梯度在特定位置出现Inf。这时可以把ignore_index单独拿出来验证一下。如果你用torch.amp.autocast(dtypetorch.bfloat16)代替FP16NaN的风险其实小很多因为BF16的范围和FP32一样。后面第5节会细说。4.4 和DDP、梯度裁剪、EMA、torch.compile组合使用时要注意什么训练不只是单卡裸奔AMP经常要和其他组件一起用这时候有几个关键姿势要记住。DDP多卡训练PyTorch的DistributedDataParallel天然支持AMP通常每个进程里各维护一个GradScaler即可。比较稳妥的做法是让scale状态在所有进程里保持一致最简单的方式是在scaler.step(optimizer)之后调用一次dist.all_reduce同步scale当然这属于进阶优化大多数情况下默认行为已经够用。梯度裁剪如果你用clip_grad_norm_或clip_grad_value_必须先调用scaler.unscale_(optimizer)把梯度的缩放还原再裁剪最后调用scaler.step(optimizer)。原因是梯度裁剪本质是对梯度做L2范数约束如果直接剪在缩放后的梯度上范数计算会完全跑偏。我见过很多新手第一步没做unscale结果梯度裁剪彻底失效。正确示例scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()EMA指数移动平均EMA一般维护FP32副本AMP下参数更新仍是在FP32上做的所以基本可以直接用。但要确保保存模型时把EMA权重转回FP32再保存否则载入到CPU或者部署环境时容易出精度问题。torch.compilePyTorch 2.x的编译模式和AMP是可以同时用的而且compile能把autocast范围内的算子进一步融合收益可能更大。但建议先单独验证一个step的loss数值没有异常再铺开整个训练否则编译错误和AMP问题混在一起排查难度会加好几倍。5. 再补充一点场景判断什么情况值得用AMP什么情况最好别动5.1 最适合AMP的场景我自己用得最多的场景是三种一是显存不够比如只有8GB或16GB显卡要跑中小型Transformer或者U-NetAMP能让batch size翻倍二是训练吞吐不够模型够大、GPU计算密度够高Tensor Core能发挥优势AMP能白赚几十分钟到几小时三是做大规模超参搜索需要快速跑很多组合AMP的显存省让调度更宽松。最近社区里经常看到“低显存运行模型”的话题很多大模型在推理阶段可以做量化但训练阶段AMP几乎是标准配置。混合精度和量化是两回事但精神本质是一样的在不明显牺牲效果的前提下用更低的数值精度换回性能和显存。5.2 数值敏感任务要格外谨慎强化学习里常用的PPO、DQN我试过直接开AMP有些场景会出现Q值估计偏差放大训练动作策略变得不稳定。原因可能是价值网络对梯度的微小扰动特别敏感AMP的舍入误差在复杂奖励信号下会被反复放大。这种任务不是不能用AMP而是需要更细致地处理比如只对部分block开启autocast或者使用BF16混合精度来替代FP16又或者把输入数据统一做标准化减少数值范围差异。如果模型本身不大、训练数据也不大只是追求一点吞吐提升那其实没必要冒险开AMPFP32更省心。5.3 再说BF16混合精度AMP的另一个选择如果你用的是Ampere架构以上的显卡比如RTX 30、40系列可以试试torch.bfloat16。BF16虽然也是2字节但指数位和FP32一致取值范围很大所以不需要Loss Scaling训练稳定性明显好很多。代码上的改动很小with autocast(cuda, dtypetorch.bfloat16): outputs model(inputs)不需要GradScaler直接loss.backward()即可。代价是精度位只有7位理论上和FP32差距比FP16大。但实践中很多大模型预训练、微调任务用BF16都非常稳U-Net这类视觉模型反而更常用FP16。你也可以自己切换对比看哪个任务收敛效果更好。我个人现在的习惯是新项目先跑一版FP32基线再看显存和速度瓶颈在哪决定用FP16还是BF16。如果模型够大、GPU也支持我会直接上BF16少操心NaN。如果显卡老一点、为了显存收益就选择FP16加GradScaler。6. 我从实际项目里总结的几句心里话AMP不是一个深不可测的黑魔法它的本质是在“精度”和“资源”之间做一次工程化的交换。我做了这几年模型训练最深的感受是工具再好也不如理解原理重要你只有知道FP16为什么丢精度、GradScaler在解决什么问题才能真正应对那些莫名其妙的NaN和Inf而不是靠幸运值在跑实验。如果你要复现这套流程建议从上面的ToyModel开始先熟悉接口用法再把AMP套进你真实的任务里。如果在显存、吞吐上遇到跟预期不符的情况优先检查是不是batch太小、GPU架构太老或者自己的自定义算子破坏了FP16路径。很多时候不是AMP没用是错误的用法让它发挥不出来。最后分享一个小技巧不管用不用AMP我都会在训练脚本里加上torch.cuda.reset_peak_memory_stats()并且每隔N个step打印一下峰值显存和每秒样本数。这种量化习惯比任何经验都更能帮你判断某个优化手段到底值不值得保留。混合精度只是一个起点真正让训练快的往往是这种不断对比和调试的过程。