☰
深度学习AMP训练必备:梯度缩放器GradScaler原理与实战排查
2026/10/5 10:52:43 网站建设 项目流程

做深度学习训练的朋友,尤其是跑过大模型、大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.12345,FP16能表示的可能是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_scale=1024甚至64都可以,但我个人建议在完全不理解任务之前先不要动,否则scale会在运行前几步频繁回退,反而影响稳定性。

接下来是它的动态调整策略,PyTorch里的默认参数是growth_factor=2.0、backoff_factor=0.5、growth_interval=2000。逻辑是:每次训练迭代中,如果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_type='cuda', dtype=torch.bfloat16)或torch.autocast(device_type='cpu', dtype=torch.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(), lr=1e-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再交给scaler:torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)scaler.step(optimizer)。这个顺序的问题是:此时梯度还是被缩放过的,你按真实阈值1.0去裁剪缩放后的梯度,等于阈值被乘了scale。如果scale已经涨到65536,实际裁掉的标准就变成了65536,几乎没有任何裁剪效果。

正确做法是:先让GradScaler把梯度unscale回真实值,再做裁剪,再step。PyTorch提供了对应API:

scaler.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的状态也存进checkpoint:

checkpoint = { "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 中断 实例”这种结果,先别急着疑惑,你撞上的是另一个AMP:Asymmetric 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的enabled=False设一下,或者直接设置init_scale=64且不调用update()。固定小scale下,梯度基本不会因为FP16溢出而报NaN,此时如果模型还能训出正常loss,说明问题出在自动缩放逻辑与某层的动态范围冲突;如果固定小scale下也训不好,那就确认模型本身有问题。这个对照实验做起来很快,能帮你砍掉一大片怀疑方向。

最后,分享一个我个人现在还在用的小习惯:断点续训时,不仅恢复GradScaler的state_dict,还会把训练日志里的scale变化画成曲线。一旦发现恢复训练的scale和中断前差距过大,我会稍微调低学习率多跑几百步,等scale自动回归到合理区间再说。混合精度训练总体来说是用一个缩放因子撬动整个FP16流程的数字稳定性,如果说模型是血肉,optimizer是心脏,那GradScaler更像一个时刻盯着血流量的自动反馈阀。把这个阀门看明白、习惯它、会用它的边界,AMP才能在给你的训练速度带来实实在在的提升,而不是变成半夜排查NaN的噩梦来源。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询