做深度学习训练的朋友尤其是跑过大模型、大batch、还嫌显存不够用的人应该都绕不过一个名词自动混合精度AMP。我第一次接触AMP是在开源检测项目里打开一行配置训练速度直接上了一个台阶但偶尔也会冒出loss变NaN、梯度消失、甚至训练曲线像心电图一样乱跳的玄学问题。这篇文章想聊的就是AMP背后那个最容易被忽视、却决定了整套数值稳定策略能否成立的关键组件——梯度缩放器GradScaler。不管你是刚开始用混合精度训练的新手还是已经踩过NaN坑、想彻底搞懂原理的进阶玩家这篇文章都能给你一些实际可用的参考。1. 先搞清楚AMP到底在解决什么问题1.1 FP16不是“省一半精度”那么简单很多人刚接触混合精度时第一反应是“把模型参数从FP32变成FP16显存减半、速度翻倍”。这个理解方向对但远不够。FP16确实是16位浮点数相比FP32的32位存储确实少了一半但它折腾的远不只是存储。FP16的数值结构是1位符号、5位指数、10位尾数FP32则是1位符号、8位指数、23位尾数。你可以简单理解为FP16用更少的位数去表示数值结果就是它能表示的数值范围窄得多能保留的有效数字也少得多。FP16的最大有限值大概是65504而FP32能表示到3.4e38这个量级。更关键的是精度FP16的有效十进制位数大约只有3到4位什么意思呢假设一个权重或者梯度是0.12345FP16能表示的可能是0.1234或者0.1235后面的两位就已经丢了。对于模型推理来说这种损失很多时候可以接受这也是为什么很多部署端模型直接用FP16、INT8做量化效果还很稳。但训练是另一回事训练过程中要通过梯度不断修正参数梯度本身就是很小的数值而且它需要积累、传递、多步迭代。如果梯度在FP16里被砍掉几位有效数字或者干脆变成0那网络就根本没有信息可以去更新了。那为什么还要用AMP因为不是所有计算都必须用FP16。AMP的核心思想是“混合”——让适合FP16的计算跑在FP16上让不适合的、对精度敏感的环节继续用FP32。比如矩阵乘法、卷积这类计算密集型的算子在硬件上有专门为FP16设计的加速单元速度明显更快而像BatchNorm里的统计运算、残差相加这类对精度要求高的操作保持FP32反而更稳。AMP就是自动帮你在模型前向和反向传播过程中按算子粒度去动态选择精度不用你手动去改模型结构。这也就是“自动”两个字的含义。真正让很多人困惑的问题是前向传播用FP16模型就老老实实算完了看起来没什么问题为什么训练时还要多出一个梯度缩放器这就得聊到训练独有的反向过程了。1.2 梯度下溢AMP落地路上最容易被忽视的坑我们反向传播计算梯度时得到的数值往往是极其小的。一批样本算下来的平均梯度常见量级在1e-3到1e-8之间甚至更小。而FP16能精确表示的规格化正数最小大约只有6.1e-5比这个更小的数值不是不能表示而是会落入非规格化区域有效数字急剧减少再小就直接变成0了。这个现象叫下溢。你可以把它理解成在一个精确到厘米的尺子上测量一根头发的直径尺子本身的刻度就不够细你量出来的结果要么是0要么是莫名其妙的一个数。训练时更麻烦的是梯度下溢为0之后那层网络的参数就完全停止更新了。表面上loss还在降但某一层已经死掉了。当然梯度下溢也不是只有AMP才有纯FP32训练里也会遇到但FP32的动态范围足够大下溢概率低得多。一旦切换到FP16这个问题会被急剧放大。既然梯度这么小那最简单的思路就是“把它放大一点再算”。于是有了梯度缩放器用一个缩放因子去乘loss再反向传播梯度就会同比放大在FP16里不再是0参数更新前再除以缩放因子恢复成真实梯度。这个思路听起来很直白但它要解决一连串工程问题放大多少合适放大后会不会导致loss或者其他数值溢出训练阶段不同梯度的量级会变化缩放因子要不要跟着调整答案就是动态调整的梯度缩放器它会在训练过程中持续监测梯度是否溢出一旦发现inf或NaN就减小缩放因子如果长时间没溢出就尝试增大缩放因子。这套机制训练过程中全自动完成所以你只需要在PyTorch里加三行代码剩下的它自己看着办。这就是为什么它叫“梯度缩放器”而不是“损失缩放器”——表面上我们缩放的是loss真正保护的其实是梯度。理解了这一点后面所有调参和排查逻辑都会顺畅很多。2. 梯度缩放器的工作逻辑与设计细节2.1 缩放、回传、还原三步闭环很多框架的梯度缩放器使用时都遵循同一个闭环逻辑我拿PyTorch的GradScaler来拆解因为它的API设计非常典型理解了之后你换成TensorFlow或者MXNet也差不多。第一步缩放前向传播后得到一个loss记为L梯度缩放器用当前缩放因子S乘以L得到缩放后的损失L_scaled L * S然后执行scaler.scale(loss).backward()。这一步目的就是让反向传播过程中计算出的梯度整体放大S倍从而避开FP16的下溢区域。需要强调的是这个乘法发生在loss计算之后、反向传播之前而且loss本身通常是一个标量。第二步回传模型自动完成反向传播得到所有参数的梯度。此时这些梯度是在放大状态下的我们不马上用它去更新参数。第三步还原与更新优化器要执行step之前缩放器先把缩放后的梯度除以S还原成真实梯度再交给优化器去更新参数。PyTorch里这个过程封装在scaler.step(optimizer)之中。它内部会先检查本轮梯度的所有元素里有没有inf或NaN如果有就认为这次更新是无效的直接跳过一步权重更新如果没有就把梯度unscale回真实值然后才调用optimizer.step()。这里有个重要的点很多人会问为什么不直接把参数和梯度都保存在FP16里那样不是更快吗如果真这么干不需要多少轮模型就会废掉。原因是更新参数用的是真实梯度而真实梯度经过还原后量级很可能是0.0001这种如果参数本身也是FP16参数值的有效数字又不够了更新量就直接丢失。所以标准的AMP实现里模型参数会维护一份FP32的“主副本”FP16计算、FP32更新。这样既享受了FP16算得快的好处又保留了FP32更新参数的精度。2.2 动态缩放初始值为什么是65536我现在还记得第一次看到GradScaler默认初始值65536时的反应怎么是这么个奇怪数字后来想明白了很有讲究。先想想缩放因子的上限约束。缩放后的loss也要在FP16能表示的范围内不然一乘就成inf后面全完了。FP16最大有限值是65504所以缩放因子和原始loss的乘积必须小于65504。初始缩放因子如果取65536那就要求原始loss大约不超过1。大多数深度学习训练任务的初始loss都在个位数甚至更小比如分类任务的交叉熵损失一般几以下回归任务的MSE也不会离谱到哪里去所以65536这个初始值在绝大多数场景够用且安全。如果某个任务的初始loss本身就很大比如直接用了带巨大数值的标签再乘上65536直接就溢出了。这时候缩放器会在第一步就检测到inf然后自动把scale缩小一半。所以初始值也不是不能动init_scale1024甚至64都可以但我个人建议在完全不理解任务之前先不要动否则scale会在运行前几步频繁回退反而影响稳定性。接下来是它的动态调整策略PyTorch里的默认参数是growth_factor2.0、backoff_factor0.5、growth_interval2000。逻辑是每次训练迭代中如果scaler.step()检测到梯度正常没有溢出就给计数器加一当连续正常步数达到2000步时就把scale乘以2。一旦任何一步检测到inf或NaN马上把scale乘以0.5然后清空计数器重新计数。这个设计很像数码相机里的自动曝光场景亮就调低曝光场景暗就调高曝光目标是把画面维持在最佳亮度范围。训练前期梯度可能比较猛scale会相对保守训练后期梯度普遍较小scale会慢慢增大防止梯度在下溢区躺平。这比固定缩放因子灵活得多也是为什么现在的AMP实现都倾向于动态缩放。2.3 什么时候其实不需要梯度缩放器不是所有混合精度都需要梯度缩放器这点很多人没搞清楚结果在BF16混合精度训练里硬套GradScaler发现反而拖慢了速度。先看BF16。BF16也是16位但它的分配是1位符号、8位指数、7位尾数指数位数和FP32完全一样。这意味着BF16的动态范围和FP32几乎相同它不会出现FP16那种梯度直接下溢为0的问题。代价是尾数只有7位精度低得可怜但在训练场景里很多硬件的BF16算子配合FP32参数更新仍然能保持收敛。PyTorch对BF16混合精度友好的设备上只需开启torch.autocast(device_typecuda, dtypetorch.bfloat16)或torch.autocast(device_typecpu, dtypetorch.bfloat16)不需要初始化GradScaler也不需要scale与unscale。如果你强行给BF16流程加缩放器反而引入无谓的乘除法速度不升反降。再看推理阶段。推理没有反向传播不存在梯度下溢问题模型参数直接以FP16读入做前向即可所以推理用的AMP通常只是autocast包裹前向过程不需要任何GradScaler。有些教程为了图省事把训练和推理代码混在一起推理时也顺手写了GradScaler这不会报错但属于无意义开销。另外还有全FP32训练以及那些算子本身不支持FP16的训练流程。这时候GradScaler完全不需要出场。判断标准很简单你的反向传播过程里是否存在FP16参与、且梯度可能小于FP16最小精确范围的环节。如果有就需要缩放器如果没有老老实实别画蛇添足。3. 实操PyTorch里把GradScaler用顺的全流程3.1 最小可用接入代码理论上一段完整的混合精度训练循环只要改动几个地方。我直接给一个最小可用示例这是最典型的PyTorch写法from torch.cuda.amp import autocast, GradScaler model Model().cuda() optimizer torch.optim.Adam(model.parameters(), lr1e-3) scaler GradScaler() for epoch in range(epochs): for inputs, labels in dataloader: inputs inputs.cuda() labels labels.cuda() optimizer.zero_grad() with autocast(): outputs model(inputs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() # 该步之后的loss.backward()、optimizer.step() # 都不能再单独出现了这段代码里有三处和普通训练不一样with autocast():包住前向传播和loss计算scaler.scale(loss).backward()替代普通的loss.backward()scaler.step(optimizer)替代普通的optimizer.step()。最后还要有scaler.update()这个调用是根据本步溢出情况更新缩放因子。我在带身边人的时候发现大家最常犯的三个错集中记录一下。第一optimizer.zero_grad()放在了autocast外面这没问题因为清零不需要前向。还有更常见的迷惑操作是把optimizer.step()放在scaler.step()后面又调了一遍等于重复更新参数loss会立刻爆掉。第二在autocast外部手动计算自定义loss比如把模型输出拿到外面做torch.softmax后再算NLL这些操作如果不在autocast上下文里会跑在FP32下理论上没什么问题但如果你在这些外部操作里用了对数值敏感的算子比如torch.log(softmax)很容易产生不精确结果。稳妥做法是所有涉及前向与loss的计算全部放进autocast上下文但权重更新除外。第三遇到loss为NaN时第一个动作是把整个AMP关掉这是最有效但也是最粗笨的排查法后面我会专门展开。3.2 梯度裁剪与梯度累积的正确处理梯度裁剪是很多模型的刚需比如NLP、强化学习不加根本训不动。但如果你的代码里混了AMP和梯度裁剪顺序错了会引起很隐蔽的bug。错误写法是先clip再交给scalertorch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)scaler.step(optimizer)。这个顺序的问题是此时梯度还是被缩放过的你按真实阈值1.0去裁剪缩放后的梯度等于阈值被乘了scale。如果scale已经涨到65536实际裁掉的标准就变成了65536几乎没有任何裁剪效果。正确做法是先让GradScaler把梯度unscale回真实值再做裁剪再step。PyTorch提供了对应APIscaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update()scaler.unscale_(optimizer)会一次性把优化器管理到的所有参数梯度除以scale。执行完之后梯度已是真实梯度clip阈值也回归直觉。再说梯度累积。梯度累积是为了在小batch上模拟大batch通常步骤是若干个mini-batch只做backward不step累计梯度冲过一定数量后再更新参数。和AMP结合的坑在于如果你对每个mini-batch都执行scaler.scale(loss).backward()梯度会被多次乘scale累积到优化器里的也是多次缩放后的梯度。最终step时GradScaler只unscale一次相当于所有小步的真实梯度相加从数学上等价于没有缩放的大batch梯度累计。这个关系是成立的所以框架层面没有大问题。但在实操中如果你的某个mini-batch产生了NaN或Inf这个脏梯度会污染整个累积窗口。scaler.step()会检测到NaN并跳过当次更新可累积窗口里的其他数值已经被污染你不得不多累积几轮。我们的经验做法是累积模式下每个mini-batch额外记录一下loss_finite torch.isfinite(loss)一旦发现某一步loss异常跳过该步的backward并在本轮step时清零优化器梯度、不执行更新。这样能把污染范围控制在单步内。3.3 分布式DDP与断点续训的细节分布式训练结合AMP很多人担心缩放因子在不同卡上不同步。实际上GradScaler是每个进程独立的它只是在本地判断本进程的梯度是否溢出。DDP的梯度all-reduce发生在backward()过程中而所有卡的scale值一致因为大家初始值一样、数据加载也一致所以all-reduce后的梯度仍然等价于真实梯度的缩放值。只要各卡之间没有数据不平衡或数值漂移尺度的一致性是能保证的。真正需要留意的是unscale_和DDP通信的先后顺序DDP在backward阶段就已经完成了梯度同步之后的unscale和step都是本地操作所以不存在顺序问题。断点续训是另一个高频需求。GradScaler内部维护了scale、growth_tracker和found_inf_per_device等信息如果你只保存模型和优化器状态恢复训练后scale会重置回65536虽然有动态调整兜底但训练前期可能产生不必要的scale波动。正确做法是把scaler的状态也存进checkpointcheckpoint { model: model.state_dict(), optimizer: optimizer.state_dict(), scaler: scaler.state_dict(), epoch: epoch, } torch.save(checkpoint, ckpt.pt) # 恢复时 ckpt torch.load(ckpt.pt) model.load_state_dict(ckpt[model]) optimizer.load_state_dict(ckpt[optimizer]) scaler.load_state_dict(ckpt[scaler])这段代码我几乎每份训练脚本里都会写。因为一旦训练了几万步后scale已经涨到几百万甚至几千万如果从65536重新开始很可能把先前的动态区间全部打乱白白损失一天训练时间。4. 常见问题与排查技巧实录4.1 NaN/Inf排查先关AMP还是先查scale训练中出现NaN是最让人抓狂的AMP只是让这个问题的定位变得复杂了一层。简单粗暴的“关掉AMP再试”确实能把混合精度因素排除掉但它无法告诉你问题是出在FP16的哪个算子还是出在模型本身。我的排查顺序是这样的。第一步把scaler.update()暂时改成不更新通过scaler.get_scale()打印出当前缩放因子看它是不是一直在减半。如果scale从65536一路掉到1024甚至更小说明系统正在持续检测到溢出。这时候可以判断大概率原始梯度本身已经包含了inf或NaN缩放器只是在拼命回退scale但你相关的模型层已经进入不可恢复的状态了。第二步使用torch.autograd.detect_anomaly()重新跑几个batch它会打印出第一次出现NaN的反向路径。这一步非常消耗性能只能小规模跑。第三步在反向传播前手动检查单层梯度比如上一个钩子打印每一层梯度norm看是哪个层先爆炸。关于NaN我还有一个很实用的经验loss出现NaN和梯度出现NaN是两回事。loss在FP32里算的可能还是有限值但backward时某一层梯度已经爆了。尤其是用了torch.exp、torch.log这类指数/对数算子的损失函数中间值溢出时并不会第一时间体现在loss上。所以排查时不要只盯loss曲线要看scaler.found_inf_per_device的状态它记录了有没有溢出以及是哪个设备挑的头。4.2 加了AMP反而更慢的三个元凶不是加了AMP就一定变快这个结论我得先说在前面。如果你的模型很小GPU没那么新很可能AMP带来的收益微乎其微甚至更慢。我自己遇到过的三个原因你可以对照一下。第一个元凶是算子不支持FP16导致频繁的FP32/FP16格式转换。比如模型里用了大量自定义Python循环、动态形状的张量操作、或者在autocast外反复做CPU同步比如.item()、.type()、.float()这类强制cast。每做一次转换都可能触发设备同步GPU流水线被卡住速度反而下降。解决办法是减少这类操作或者在显式支持AMP的算子体系里重写。第二个元凶是梯度缩放器本身带来的开销。对于BF16设备缩放乘除就是纯消耗前面说过不要给BF16加GradScaler。对于FP16设备如果每次backward后又额外打印梯度、强行unscale所有参数这些开销在某些小模型上甚至会抵消FP16的提速。我的建议是先用torch.autocast跑一把不带GradScaler的速度再用带GradScaler的速度对比这样能明确量出缩放器带来的开销占比。第三个元凶是CPU瓶颈。AMP主要加快GPU计算如果你的GPU利用本来就不高数据加载和预处理是瓶颈那AMP提速的意义就很小。此时应该先看nvidia-smi里的GPU利用率不到80%说明瓶颈在别处。优化数据管线比如用num_workers、用pin_memory比纠结AMP参数管用得多。4.3 搜“amp”时撞见的另一个世界rk3506与嵌入式多核中断如果你是因为想查AMP相关资料结果搜出一堆“rk3506 amp 中断 实例”这种结果先别急着疑惑你撞上的是另一个AMPAsymmetric Multi-Processing非对称多处理常见于嵌入式多核芯片架构里。这里面的“amp”指一个芯片上不同的处理器核心跑不同的角色、执行完全不同的任务比如一个高性能核心负责应用计算另一个低功耗核心负责实时控制核心之间用mailbox中断来通信和协作。它跟深度学习的自动混合精度除了缩写撞车没有半点关系。我在RK系列芯片平台上调过一段时间的多核通信中断这个区分让我印象很深。嵌入式领域的AMP重点在于中断处理和核间通信的确定性任何在中断上下文里执行耗时浮点计算、动态内存分配、甚至打印日志的行为都可能破坏实时性。我记得有一次做核间通信测试把一段浮点矩阵运算误放进中断处理函数里结果系统响应时间从几十微秒直接飙到毫秒级几个控制任务全部超时。后来把计算任务挪到非实时核心的普通线程里中断只负责置标志位、搬数据系统才恢复稳定。这一点放到AMP自动混合精度的语境下也值得一提如果你在边缘设备上做深度学习训练或推理同样不要在中断线程里做浮点重负载操作更不要把GradScaler这类有状态更新逻辑的东西放进中断钩子里。混合精度训练和推算是普通线程里的活中断只需要做最小化的事件通知。搜索时遇到“amp 中断”这类词先确认一下到底是指哪个AMP不然很容易被领域黑话带偏。5. 我的几个使用习惯以及不建议模仿的骚操作5.1 每个新项目必做的固定动作我现在每开一个新的训练任务无论模型多简单都会主动加上一套AMP验证流程这对排查问题很有帮助。固定动作一是先跑300步的纯FP32基线记录loss曲线和吞吐量。然后开AMP同样跑300步对比两条loss曲线。如果在同样的学习率下loss曲线明显变抖或者scale一路掉说明要么模型里有对精度极敏感的稀疏层要么学习率本身已经太高。这个对比很朴素但能快速暴露一半以上的问题。固定动作二是每个epoch记录一次当前scale值。不要只关心get_scale()返回的数还要看它的变化趋势。如果scale连续几百步都停在同一个数值不动说明梯度一直在下溢区域里打转此时模型可能已经“假死”了loss虽然还在变但某些层已经没有任何有效更新。遇到这种情况优先检查是否有Sigmoid、Tanh这类有饱和区的激活函数以及初始化的标准差是不是太小。5.2 不要自己造轮子也别忘了关自动缩放试一把有人喜欢自己写动态缩放逻辑比如“如果loss大了就除以2小了就乘以2”这种土办法是危险的。loss大小和梯度溢出没有直接对应关系你用loss做判断等于在错误的信号上做反馈控制。框架内置的GradScaler通过检测inf/NaN来判断溢出才能准确捕捉到FP16下溢问题。所以我的态度很明确能做轮子但不要把时间花在重造这种已经被验证过的轮子上。但反过来也建议你在排查时把自动缩放关掉换成固定scale跑一把。方法很简单把GradScaler的enabledFalse设一下或者直接设置init_scale64且不调用update()。固定小scale下梯度基本不会因为FP16溢出而报NaN此时如果模型还能训出正常loss说明问题出在自动缩放逻辑与某层的动态范围冲突如果固定小scale下也训不好那就确认模型本身有问题。这个对照实验做起来很快能帮你砍掉一大片怀疑方向。最后分享一个我个人现在还在用的小习惯断点续训时不仅恢复GradScaler的state_dict还会把训练日志里的scale变化画成曲线。一旦发现恢复训练的scale和中断前差距过大我会稍微调低学习率多跑几百步等scale自动回归到合理区间再说。混合精度训练总体来说是用一个缩放因子撬动整个FP16流程的数字稳定性如果说模型是血肉optimizer是心脏那GradScaler更像一个时刻盯着血流量的自动反馈阀。把这个阀门看明白、习惯它、会用它的边界AMP才能在给你的训练速度带来实实在在的提升而不是变成半夜排查NaN的噩梦来源。