返回专辑
·Johan·18 分钟阅读

混合精度下的梯度累积:loss 缩放、裁剪与 step 顺序

沿 autocast、GradScaler、unscale_、裁剪与 optimizer.step 的数据流说明静默错误;区分 micro-batch 平均与求和,并用等价性测试验收。

混合精度下的梯度累积:loss 缩放、裁剪与 step 顺序

显存不够时上梯度累积,再叠 AMP,训练「能跑」却对不齐基线,是很常见的静默事故。根因通常不是模型结构,而是缩放、反缩放、裁剪、清梯度与 step 的顺序在 micro-batch 边界上错位:有的实现每个 micro-batch 都 scaler.step,有的在 unscale_ 之前裁剪,有的把累积平均和 loss_scale 缠在一起。

主线:先固定数学上的有效 batch,再固定 AMP 状态机(含 GradScaler 内部),对照 DeepSpeed / FSDP / Hugging Face Trainer 的差异,最后用数值例子与等价性单测锁死实现。

1. 梯度累积在数学上在干什么

设目标是大批次损失

但显存一次只能吃 条,。把 拆成 个 micro-batch:

对参数

实现上有两种等价写法:

  • 先平均:每次 loss = L_k / Kbackward 累积;
  • 先求和再统一除:每次对 backward,step 前把梯度除以

与 DDP 联用时,还要关掉前 次的 all-reduce(no_sync),只在最后一次同步。漏了 no_sync 不会直接 NaN,但通信量和有效梯度语义都会偏:每次 micro-batch 都做一次平均,等价于用错误频率的「半同步」更新。

数值上先钉死一件事:若 ,有效 batch 是 。日志里写 batch_size=8 却按 调学习率,或反过来,都会让「AMP 有没有坏」根本无法判断——基线本身就在飘。

2. AMP 状态机:scale 的是什么

PyTorch AMP 的典型路径:

  1. autocast 下算 loss
  2. scaler.scale(loss).backward() 把损失乘以 loss_scale 再反传,使 fp16 梯度更不易 underflow;
  3. scaler.unscale_(optimizer).grad 除以同一 scale,得到真实梯度;
  4. 可选:梯度裁剪(必须在真实梯度空间);
  5. scaler.step(optimizer):若检测到无效梯度则跳过 step,并调整 scale;
  6. scaler.update()

关键不变量:裁剪与任何基于梯度范数的决策,都必须在 unscale_ 之后。在 scaled 空间裁剪,等价于把阈值乘上了未知的 loss_scale,有效更新会被无声扭曲。

3. GradScaler 内部在维护什么

GradScaler 不是「把 loss 乘个常数」的语法糖,而是一套带滞后决策的状态机。理解它,才能解释为什么「每个 micro-batch 都 update」会把曲线抽坏。

记当前缩放因子为 。对一次优化 step:

  1. scale(loss) 返回 ;backward 后 .grad 处于 scaled 空间;
  2. unscale_ 就地除回,并扫描是否存在 inf/nan
  3. 若发现无效梯度:step 不调用 optimizer.step,并把「本步失败」记入内部计数;
  4. update() 根据连续成功/失败步数调整 :连续失败则 ,连续成功满阈值则

默认量级上,初始 常取 一类大正数;backoff 因子常见为 ,growth 间隔常见为 次成功 step——具体以所用 PyTorch 版本文档为准,验收时应用 scaler.get_scale() 打点,不要背死硬编码。

对梯度累积,一次「优化 step」对应 次 backward、一次 unscale、一次 step、一次 update。若在每个 micro-batch 后都 update,等于把「是否 overflow」的采样频率放大了 倍:某一个 micro-batch 的偶发尖峰会立刻打掉 ,而合法的其余 个 micro-batch 从未有机会组成完整更新。结果是:训练仍在跑,但实际更新次数与学习率调度脱节,loss_scale 曲线呈锯齿。

另一个内部细节:unscale_step 之间不要再对同一 optimizer 做第二次 unscale_。GradScaler 用内部标志防止重复反缩放;重复调用会直接报错或(在错误封装下) silently 跳过——两种都比「看起来训了一会儿」更糟,因为 CI 若没断言会漏掉。

python
# 概念示意:一次优化步内 scaler 的契约
# found_inf = any(~isfinite(p.grad) for p in params after unscale)
# if found_inf: skip optimizer.step(); request smaller S
# else: optimizer.step(); maybe grow S after N successes

4. 累积场景下的正确骨架

python
optimizer.zero_grad(set_to_none=True)
for micro_step, (x, y) in enumerate(micro_batches, start=1):
    with torch.cuda.amp.autocast(dtype=torch.float16):
        loss = criterion(model(x), y) / accumulation_steps
    scaler.scale(loss).backward()

    if micro_step % accumulation_steps == 0:
        scaler.unscale_(optimizer)
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
        scaler.step(optimizer)
        scaler.update()
        optimizer.zero_grad(set_to_none=True)

几点刻意为之:

  • 除以 accumulation_steps 放在 scale 之前(对 平均);scale 只负责 AMP,不负责累积语义;
  • 每个优化 step 只 unscale_/step/update 一次;
  • zero_grad 在 step 之后,避免清空未完成的累积。

错误示范 A:每个 micro-batch 都 scaler.step——等价于用更小 batch 更频繁更新,且 scale 动态被错误驱动。
错误示范 B:累积完直接 clip_grad_norm_scaler.step,忘记 unscale_——裁剪阈值被 scale 污染。
错误示范 C:loss = scaler.scale(L_k) / Kscaler.scale(L_k / K) 混用时,若再手动除梯度,会双重归一化。

数值例子(便于 code review 对照)。设某参数在三个 micro-batch 上的真实梯度分别为 ,当前

  • 正确「先平均」:每次 backward 贡献 ,累积后 .gradunscale_ 后得 ,等于
  • 错误「忘记除以 」:累积 scaled 梯度 ,unscale 后得 ,有效学习率被放大 倍。
  • 错误「在 scaled 空间按 max_norm=1 裁剪」:范数先按 量级裁到 ,再 unscale 得到 ,更新几乎被抹掉。

把这组数字写进单测的期望值,比口头说「别忘了除以 K」更能挡住回归。

5. bf16 与 fp16 的差别不要抄错

bfloat16 指数范围更宽,很多模型可以不用 GradScaler。若关闭 scaler,就不要残留 unscale_ 路径;若仍使用 scaler,行为应与文档一致,避免「半开 AMP」。同一套累积代码在 fp16/bf16 间切换时,用标志位显式分支,比依赖默认更安全。

GradScaler 在连续 overflow 时会缩小 scale;若每个 micro-batch 都 update,scale 会被错误频率的 overflow 信号抽打。这是「能训但曲线怪」的常见来源。

实践上建议:

  • A100 / H100 上优先试 bf16 + 无 scaler 的累积骨架;
  • 仍用 fp16 时,把 scaler.get_scale() 与「每 N 个优化 step 的跳过次数」写进日志;
  • 切换精度时复跑第 10 节的等价性测试,而不是只看第一个 epoch 的 loss。

6. 与梯度裁剪、权重衰减、梯度惩罚的交互

  • 裁剪:对累积后的总梯度做一次。micro-batch 级裁剪会改变有效更新,且与大批次基线不对齐。
  • AdamW / weight decay:解耦衰减在 optimizer.step 里;只要不在 scaled 梯度上 step,语义通常保持。仍建议用等价性测试确认。
  • 不同参数组学习率unscale_ 按 optimizer 持有的参数来;确保所有该更新的参数都挂在同一个被 unscale 的 optimizer 上,或对每个 optimizer 调用 unscale_

梯度惩罚(gradient penalty) 与累积叠在一起时,最容易静默漂正则强度。以 WGAN-GP 风格为例,惩罚项形如

若主损失按 平均,但 每个 micro-batch 全量 backward、却只在第 步更新,则有效 被放大约 倍;若惩罚只在最后一个 micro-batch 计算,又会与「真大批次上对全部样本估惩罚」不对齐。可验收的写法是:把惩罚放进与主损失同一套平均约定——每个 micro-batch 计算局部惩罚并除以 ,或显式文档化为「每优化 step 估一次惩罚」并相应调整

AMP 下算 时还要注意:惩罚依赖二阶信息路径,部分算子在 autocast 下不稳定。工程上常见做法是惩罚路径强制 fp32(with autocast(enabled=False)),主损失仍走 fp16/bf16;此时 不要 对已经是 fp32 的惩罚 loss 再套一层错误理解的 scale 语义——应统一 scaler.scale(total_loss).backward(),让 scaler 只看到标量总损失。

python
with autocast(dtype=torch.float16):
    main = main_loss(model(x), y)
with autocast(enabled=False):
    gp = gradient_penalty(disc, x_real.float(), x_fake.float())
loss = (main + gp) / accumulation_steps
scaler.scale(loss).backward()

7. DeepSpeed / FSDP 与手写 AMP 累积的差异

手写 GradScaler 路径假设:梯度落在 optimizer 管理的参数 .grad 上,由你显式 unscale / clip / step。DeepSpeed 与 FSDP 把其中几段收进引擎,抄手写骨架会双重缩放或裁剪两次。

DeepSpeed。 ZeRO 阶段会切分优化器状态与梯度;deepspeed.initialize 之后通常走 model_engine.backward(loss)model_engine.step()。loss scaling 由配置里的 fp16.loss_scale / 动态 loss scale 管理,而不是你再包一层 torch.cuda.amp.GradScaler。梯度累积在配置项 gradient_accumulation_steps 中声明,引擎在内部决定何时 step。若在 DeepSpeed 外再手动 scaler.scale,等于把 乘了两次。裁剪同样应走 DeepSpeed 的 gradient_clipping 配置,避免在未汇齐的分片梯度上按全局范数裁剪。

FSDP。 全分片参数在 backward 后梯度亦分片;clip_grad_norm_ 需要能处理分片语义的版本(PyTorch 为 FSDP 提供了相应路径)。AMP 可用 torch.cuda.amptorch.amp,但 unscale 与 step 仍须落在「完整优化步」边界。FSDP + 累积时,前 次 backward 仍应避免多余通信;具体 API 随版本在 no_sync 类上下文上演进,升级框架后要重跑通信次数探针。

对照表(实现审查用):

路径谁做 loss scale谁在何时 step裁剪落点
手写 DDP + GradScalerGradScaler你在 次后显式 scaler.stepunscale_ 后、参数 .grad
DeepSpeed ZeRO + fp16引擎配置gradient_accumulation_steps 到齐后 engine.step引擎 gradient_clipping
FSDP + torch.ampGradScaler 或 bf16 无 scaler你在累积边界 stepFSDP 感知的 clip_grad_norm_

选型建议:单机/简单 DDP 用手写骨架最透明;多机大模型优先让 DeepSpeed/FSDP 管缩放与裁剪,不要混用两套状态机。迁移时用同一小模型对比「手写 vs 引擎」的前若干 step 参数差,而不是直接上全量数据。

8. 多优化器、参数冻结与「半个 unscale」

检测任务里常见:骨干一个 optimizer,新头另一个;或 LoRA 只更新少量参数。GradScaler.unscale_ 以传入的 optimizer 为准,只处理该 optimizer 参数组里的 .grad。若冻结参数仍被误加入 optimizer,或可训练参数分散在两个 scaler 状态之外的 optimizer 上,会出现:

  • A 优化器被 unscale + step,B 仍停在 scaled 梯度上就被 step
  • 裁剪只覆盖部分参数,全局范数失真;
  • 冻结参数上残留 scaled .grad,下次解冻时「炸一下」。

约定应写死:每个参与 AMP 更新的 optimizer 都走同一套 unscale → clip(可选)→ step;冻结参数 requires_grad=False 且不进 optimizer;累积边界上禁止只更新其中一组却清空另一组已累积的梯度。单测可构造两套 nn.Linear、两个 SGD,故意漏掉一个 unscale_,断言参数更新比率偏离

9. Hugging Face Trainer 里的典型陷阱

Transformers 的 Trainer 把累积、AMP、裁剪、调度器缠在同一循环里,省事也藏错。结合源码阅读与事故复盘,下面几类最常见:

  1. gradient_accumulation_steps 与全局 batch 心算不一致。
    有效 batch per_device_train_batch_size × GPU数 × gradient_accumulation_steps(再乘上 gradient checkpoint 无关的并行度)。只改累积步数却不改学习率/调度,会把「AMP 问题」误判成「优化器问题」。

  2. 自定义 compute_loss 又除一次
    Trainer 在累积场景通常已对 loss 做了平均约定(随版本可能是 /gradient_accumulation_steps 或等价逻辑)。若在 compute_loss 里再除一次,梯度系统性偏小。验收:打日志打印 loss.detach() 与反传后某层 grad.abs().mean(),对照关闭累积的 run。

  3. 回调里手写 scaler 或额外 optimizer.step
    与 Trainer 内部 step 打架,导致 double step 或在 micro-batch 边界 step。自定义训练循环时,要么完全接管,要么只读不写优化器状态。

  4. fp16bf16 同时开启或与 DeepSpeed 配置重复。
    一处开 TrainingArguments.fp16=True,另一处 DeepSpeed json 再开动态 loss scale,审查时要当「单一真源」处理。

  5. max_grad_norm 以为失效。
    实际是裁剪发生在 Trainer 管理的 unscale 之后,但你在回调里用未 unscale 的 grad 画图,误以为范数「总是爆」。探针必须挂在与 Trainer 相同的生命周期点。

  6. 日志 loss 是 micro-batch 还是优化 step 平均。
    对比基线前先统一口径,否则「AMP 让 loss 变了」可能只是记录窗口变了。

最小复现策略:用 Trainer 跑一个 toy 回归,gradient_accumulation_steps=4,再写一个纯 PyTorch 循环对齐同一数据与初始化;参数 L2 差应在 AMP 容差内。这比在万亿 token 任务上调参便宜得多。

10. 再谈平均因子:与 reduction、DDP、样本权重纠缠时

CrossEntropyLoss(reduction="mean") 已经在 micro-batch 内做了 平均。累积时再除以 ,对应的是「对全局 个样本的均值损失」。若某条路径改成 reduction="sum",却仍除以 而不除以 ,有效梯度会被放大约 倍——这类错误在代码审查里极难看出来,因为「除以 accumulation_steps」看起来是对的,错在 reduction 契约变了。

带 sample weight / mask 时更绕。设有效样本数

若各 micro-batch 的 差异大(padding、句子长度、类别不平衡采样),简单除以 并不等于对全局加权均值。两种可辩护的约定:

  • 按 micro-batch 等权:接受与真大批次加权均值的偏差,文档写清;
  • 按样本权重汇总量:backward scale(w_sum_k / w_sum_total * L_k) 或等价形式,使 贡献与一次性吃满数据一致。

DDP 下还有第三层:all_reduce 平均的是各卡梯度。若每卡累积 步且最后一次同步,语义是「全局有效 batch = 」。有人在 loss 上再除以 world_size,等于把数据并行的平均做了两次。对照实验必须固定:reduction、、world size、是否 no_sync

11. 静默错误如何表现为「几乎正确」

这类 bug 很少直接炸:

  • loss 平滑下降,但同一 epoch 的训练 loss 比全精度大批次基线高一截;
  • 学习率扫描最优值整体偏移;
  • 偶发 step 被 scaler 跳过,实际更新次数变少,却仍按原 epoch 调度衰减;
  • 范数裁剪几乎总是触发或几乎从不触发——阈值落在错误量纲;
  • get_scale() 呈高频锯齿,与 overflow 真实发生率不符;
  • 打开累积后 Adam 的二阶矩估计「跟上了更大噪声」,但实际是梯度少除了
  • 同一组超参在关闭累积后突然发散或收敛变慢——往往是此前「错的 」碰巧补偿了学习率。

所以要用参数更新等价性而不是「loss 是否下降」做验收。

12. 动态 loss scale 与跳过步对调度器的二次伤害

GradScaler 跳过非法 optimizer.step 时,参数没更新,但许多项目仍让学习率调度器前进。于是出现:名义上 10k step 的 cosine,实际有效更新只有 9.6k,且跳过密集发生在训练前期(scale 尚在搜索)。更隐蔽的是:跳过步若不清空 .grad,下一轮累积会把失败 micro-batch 的 inf 残渣带进合法区域——官方推荐是失败后 optimizer.zero_grad,与成功路径的清梯度对齐。

建议在日志中同时记录:

  • scale 时间序列;
  • skipped_steps / optimizer_steps
  • scheduler_steps 是否与成功 optimizer.step 一致。

若跳过率持续高于极低水平(例如长时间高于几个百分点),优先查数值爆炸点(学习率、初始化、坏数据),而不是把 growth_interval 调到极端去「硬撑」fp16。

13. 验收:等价性单测、溢出注入与数值探针

等价性。 固定种子、固定数据顺序、关闭非确定性 cudnn,在极小模型上比较:

  1. 真大批次 fp32、无累积;
  2. 累积 步的 fp32;
  3. 累积 步的 fp16+scaler。

在若干 step 内比较参数差: 应接近机器噪声; 允许 AMP 误差但方向应一致,且不出现系统性尺度偏移(例如总是约小 倍或小 loss_scale 倍)。

完整单测伪代码(可直接改成 pytest):

python
def test_amp_grad_accum_matches_full_batch():
    torch.manual_seed(0)
    K, b, dim = 4, 8, 16
    x = torch.randn(K * b, dim)
    y = torch.randn(K * b, 1)

    def run(use_amp: bool, accum: int):
        model = nn.Linear(dim, 1)
        opt = torch.optim.SGD(model.parameters(), lr=0.05)
        scaler = GradScaler(enabled=use_amp)
        # 拷贝同一份初始权重到对比组……
        opt.zero_grad(set_to_none=True)
        for k in range(accum):
            xb = x[k * b : (k + 1) * b]
            yb = y[k * b : (k + 1) * b]
            with autocast(enabled=use_amp):
                loss = F.mse_loss(model(xb), yb) / accum
            if use_amp:
                scaler.scale(loss).backward()
            else:
                loss.backward()
        if use_amp:
            scaler.unscale_(opt)
            scaler.step(opt)
            scaler.update()
        else:
            opt.step()
        return model.weight.detach().clone()

    # 基线:一次吃满 K*b(fp32)
    base = run_full_batch_fp32(x, y)
    w_fp32_accum = run(use_amp=False, accum=K)
    w_fp16_accum = run(use_amp=True, accum=K)
    assert torch.allclose(base, w_fp32_accum, atol=1e-5, rtol=1e-4)
    # AMP 允许更大 atol,但禁止 ~1/K 或 ~1/scale 的系统性偏差
    rel = (w_fp16_accum - base).abs() / base.abs().clamp_min(1e-6)
    assert rel.median() < 0.05
    assert not torch.allclose(w_fp16_accum, base / K, atol=1e-3)

def test_clip_after_unscale_unit():
    # 构造已知 grad,scale=8,max_norm=1
    # 断言 unscale 后裁剪,最终 grad 范数 ~1
    # 断言「先裁剪再假想 unscale」路径与期望相差 factor==scale
    ...

def test_overflow_skips_only_optimizer_step():
    # micro-batch 2 注入 inf loss;断言该优化步不更新参数
    # 断言 scaler.get_scale() 下降一次而非 K 次
    ...

溢出注入。 人为制造某一 micro-batch 的巨大 loss,确认:

  • scaler.step 跳过非法更新;
  • 合法累积不被半更新污染(通常需要在非法 step 清空梯度并保持与官方推荐一致的行为);
  • loss_scale 按预期下降,而不是每个 micro-batch 抽打一次。

范数探针。unscale_ 后记录 grad_norm;对比「错误地在 scale 前记录」应差一个 loss_scale 因子。把该探针开在 CI 的单测里,能挡住回归。

调度器对齐。 断言 scheduler.last_epoch(或等价计数)按优化 step 递增:跑 个 micro-batch、K=4,调度器应只走 步。

通信次数探针(DDP)。 在假参数上挂钩子,统计一次优化步内的 all-reduce 次数:期望为 (或与桶数相关的固定倍数),而不是 。这比看 NCCL 日志更适合放进单测。

14. 现场排查路径(从症状到探针)

遇到「开了累积 + AMP 后对不齐」时,按下面顺序收敛,避免同时改学习率、改 、换精度:

  1. 关掉 AMP,只保留累积,对齐 fp32 真大批次。若此步失败,问题在平均因子 / no_sync / reduction,与 scaler 无关。
  2. 打开 bf16 且无 scaler,再对齐。若 bf16 正常、fp16+scaler 失败,焦点转到 unscale/clip/update 顺序。
  3. 检查是否双重缩放:搜索代码与配置里所有 GradScalerloss_scalefp16 开关,只留一处。
  4. 核对裁剪量纲:日志同时打印 unscale 前、后的 grad_norm,比值应接近当前 get_scale()
  5. 核对跳过步:若 fp16 跳过频繁,先找坏 batch / 过大学习率,再谈累积。

上次在分割模型上排查时,fp32 累积已对齐,fp16 却系统性偏小约一个数量级——根因是裁剪写在 unscale_ 前,而当时 loss_scale 恰好在 量级附近游荡。探针一把比值打出来,争论立刻结束。另一例是 Trainer 自定义 compute_loss 重复除以累积步数:等价性测试在纯 PyTorch 循环里先绿,接入 Trainer 后立刻红,定位时间从「猜 AMP」缩短到「比一行除法」。

15. 与 activation checkpointing、累积窗口的边界

梯度检查点(activation checkpointing)省显存的方式是重算前向,不改变「有效 batch」数学;但它会改变一次 micro-batch 的算力与峰值显存形状。和累积叠用时,常见误判是:以为「再加大 」总是免费的。实际上 增大延长了参数更新间隔,BatchNorm / 某些依赖 batch 统计的层在 micro-batch 很小时统计更噪;若又开了同步 BN,还可能和 no_sync 的意图打架。

可操作的约束:

  • 优先减小 micro-batch 内的激活峰值(checkpoint、更小序列切分),再增大
  • 对 BN 类模型,要么固定较大的 ,要么换成 GroupNorm / LayerNorm,避免「有效 batch 只存在于纸面」;
  • 记录「优化 step 间隔墙钟时间」:若 大到调度器与正则(如 dropout 期望)语义变味,应回到真大 batch 或换并行策略。

AMP 下 checkpoint 重算的前向也必须落在同一 autocast 上下文,否则重算段落变成 fp32、主段落 fp16,数值与性能都会怪。这不是累积独有,但累积加长了一步内的重算次数,更容易在 profiling 里被当成「scaler 有病」。

16. 完整训练步的状态机(可贴进设计文档)

把一个优化步写成显式状态,比口口相传的「注意顺序」更抗回归:

  1. EnterStepzero_grad 已在上一步末尾完成;scale=S 只读。
  2. MicroForwardautocast 前向;构造 loss_k = L_k / K(或文档化的加权形式)。
  3. MicroBackwardscaler.scale(loss_k).backward();若 DDP 且 ,处于 no_sync
  4. MaybeOverflowSignal:仅在 后进入下一步;中途不 update
  5. Unscaleunscale_(optimizer);记录 grad_norm_raw
  6. Clipclip_grad_norm_;记录 grad_norm_clipped
  7. StepOrSkipscaler.step;若 skip,则 zero_grad 并打点 skipped+=1,调度器按约定是否前进。
  8. UpdateScalescaler.update() 一次。
  9. ExitStepzero_gradscheduler.step(若绑定成功更新);写日志。

代码审查对照此表逐项打勾,比只搜有没有 unscale_ 更能发现「回调里偷偷 step」类问题。

17. 学习率与有效 batch 的换算不要混进 AMP 锅里

累积把有效 batch 从 拉到 时,线性缩放学习率(或等效的 Adam epsilon / warmup)是优化问题,不是数值问题。若在错误的平均因子下「碰巧」用偏小的学习率训出能看的曲线,修复平均因子后必须重做学习率扫描,否则会把正确实现误判为「AMP 让模型变差」。验收顺序应固定为:先锁死梯度等价,再谈学习率与正则强度。把这两步拆开,能少掉大量跨团队扯皮。

同样地,不要用「把 max_grad_norm 调小」去掩盖少除了 的梯度——范数阈值会被当成新的超参债,换数据集后立刻失效。凡是能被等价性单测一次性钉死的问题,就不要留到超参搜索里碰运气。

18. 实施清单

  • 明确文档:有效 batch 、micro-batch 、累积次数 ,以及 loss reduction / 样本权重约定;
  • 代码审查焦点:scale → backward → (累积结束) → unscale → clip → step → update → zero_grad
  • DDP 下确认 no_sync 包裹前 次,并用通信次数探针锁住;
  • 调度器按成功的优化 step计数,不按 micro-batch 计数,失败 step 与清梯度策略写清;
  • DeepSpeed/FSDP/Trainer 只保留一套 loss-scale 状态机;
  • 梯度惩罚与主损失共享同一平均因子,或显式重标定
  • checkpoint 与 autocast 边界一致;BN/归一化层与小 micro-batch 的相容性写进选型说明;
  • 多优化器各自完成 unscale/step,冻结参数不进状态机;
  • 基线对比失败时先查平均因子,再查 AMP 顺序,最后查框架是否双重缩放。

混合精度解决的是数值范围与吞吐;梯度累积解决的是等效 batch。两者叠加时,任何「顺便除一下」都可能同时碰坏两套语义。把顺序写成状态机,把等价性写成测试,训练曲线上的玄学才会变成可定位的实现错误。

19. 收束前的一条硬规则

若等价性测试与通信探针未进 CI,就把「已支持 AMP 累积」从发布说明里删掉。口头保证在换 PyTorch 小版本、换 Trainer、换 ZeRO 阶段时一律不可信;能自动红的测试才是状态机的一部分。

对评审者而言,只需追问三件事:有效 batch 如何定义;unscale 与裁剪谁先谁后;失败 step 是否与调度器和解。三问都能指向代码行号与单测名,就把这三问写进合并清单,作为硬门槛——工程目标才算落地。

← 全部文章

johan's blog