跳到主要内容

数据截至 (上游 commit fd01e35c83d8)

04 · 混合精度与梯度累积

这一章讲什么: Accelerator(mixed_precision="fp16", gradient_accumulation_steps=8) 这两个开关背后,forward、backward、optimizer.step、scheduler.step 各自被改成了什么样。它们共用同一个节拍器 GradientState,是 Accelerate「包装哲学」最集中的一次展示。


1. 它要解决的小问题

混合精度训练的样板代码(手动包 autocast、scaler 缩放 loss、溢出时跳 step、步进 scheduler 前查溢出)又臭又长,且因设备/后端而异。梯度累积的样板同样烦人:loss 要除以累积步数、DDP 下中间步要用 no_sync 省掉 all-reduce、optimizer 与 scheduler 要 N 步才走一次。手写这些,换个后端就得重调。

2. 直觉

Accelerate 把两件事都收敛成一处状态 + 三层包装

  • 一处状态:GradientState.sync_gradients——「本步是否要真的更新参数」。
  • 三层包装:model 的 forward 被包进 autocast(混合精度);optimizer 的 step 被门控(不同步就空转,有 scaler 就走 scaler 并查溢出);scheduler 的 step 被门控(optimizer 没走我不走)。

accelerator.backward(loss) 则是这套机关的点火开关:先帮你 loss / accumulation_steps,再按后端选择 scaler.scale(loss).backward() 还是裸 backward。

3. 图示

一个累积周期(N=4)里各对象的行为:

step: 1 2 3 4 (同步步)
│ │ │ │
sync_gradients: False False False True
DDP: no_sync ────────────────────────▶ 正常 all-reduce
backward: scaler.scale(loss/4).backward()(每步都做)
optimizer: 空转 空转 空转 scaler.step → update
scheduler: 冻结(_step_count 补偿) → step 一次

混合精度的 forward/backward 链路:

用户 model.forward ──▶ autocast(fp16/bf16) ──▶ 原 forward ──▶ convert_outputs_to_fp32
loss ──▶ accelerator.backward: loss/N ──▶ scaler.scale(loss).backward()
optimizer.step ──▶ scaler.step(optimizer) ──▶ scaler.update()(溢出则整步跳过)

4. 原理演示(示意代码)

class Accelerator:
def backward(self, loss):
loss = loss / self.gradient_accumulation_steps # 累积归一
if self.scaler is not None:
self.scaler.scale(loss).backward() # fp16:缩放后反传
else:
loss.backward()

@contextmanager
def accumulate(self, *models):
self._do_sync() # step += 1; sync_gradients = (step % N == 0)
ctx = nullcontext() if self.sync_gradients else self.no_sync(models[0])
with ctx: # 非同步步:DDP 不 all-reduce
yield

class AcceleratedOptimizer:
def step(self):
if not self.gradient_state.sync_gradients:
return # 累积中间步:空转
if self.scaler is not None:
self.scaler.step(self.optimizer) # scaler 内部溢出则跳过真 step
self.scaler.update()
else:
self.optimizer.step()

class AcceleratedScheduler:
def step(self):
if not self.gradient_state.sync_gradients:
return # 跟着 optimizer 冻结
if any(opt.step_was_skipped for opt in self.optimizers):
return # fp16 溢出:optimizer 没走,我也不走
self.scheduler.step()

5. 真实实现

5.1 开关落地:native_amp 与 scaler 的创建

Accelerator.__init__src/accelerate/accelerator.py:562-613 决策:

  • fp16(非 CPU、非 DeepSpeed/Megatron):native_amp = True,并按后端选 scaler——FSDP2 用 get_fsdp2_grad_scaler,其余走 get_grad_scalersrc/accelerate/accelerator.py:578-583);
  • bf16:native_amp 取决于 is_bf16_available(CPU/XPU/HPU 直接为真),不建 scaler——bf16 动态范围与 fp32 相同,无需 loss scaling(src/accelerate/accelerator.py:585-597);
  • fp8:恒 native_amp = True,MS-AMP 后端即便 bf16 autocast 也要 scaler(src/accelerate/accelerator.py:601-612)。

5.2 forward 的改写链

prepare_model 里(src/accelerate/accelerator.py:1818-1827):

model._original_forward = model.forward
autocast_context = get_mixed_precision_context_manager(self.native_amp, self.autocast_handler)
model.forward = convert_outputs_to_fp32(autocast_context(model_forward_func))

get_mixed_precision_context_managersrc/accelerate/utils/modeling.py:2075-2110)按 state.mixed_precision 返回 torch.autocast(device_type=..., dtype=float16/bfloat16),native_amp 为假时返回 nullcontext()——所以无混合精度时这层包装零开销。convert_outputs_to_fp32src/accelerate/utils/operations.py:939-947)在 forward 出口把输出转回 fp32,保证 loss 计算在高精度下进行。

5.3 backward:缩放与分派

Accelerator.backwardsrc/accelerate/accelerator.py:2818-2849)的顺序:

  1. 非 DeepSpeed 时 loss = loss / self.gradient_accumulation_stepssrc/accelerate/accelerator.py:2837-2839)——DeepSpeed 的 engine 自己在 backward 里做这一步,故跳过;
  2. DeepSpeed → deepspeed_engine_wrapped.backward(loss, sync_gradients=...)
  3. 有 scaler → self.scaler.scale(loss).backward(**kwargs)src/accelerate/accelerator.py:2845-2846);
  4. LOMO optimizer 走 lomo_backward;其余裸 loss.backward()

配套的 unscale_gradientssrc/accelerate/accelerator.py:2911)与 clip_grad_norm_src/accelerate/accelerator.py:2946)处理 fp16 下的剪梯度:先 unscale 再剪,DeepSpeed/FSDP 各有特判。

5.4 节拍器:GradientState 与 _do_sync

GradientStatesrc/accelerate/state.py:1231)是第三个 Borg 单例,持有 sync_gradients、累积插件参数(num_steps/adjust_scheduler/sync_with_dataloader,来自 GradientAccumulationPluginsrc/accelerate/utils/dataclasses.py:981-1024),以及 dataloader 尾部状态(见第 3 章)。

Accelerator._do_syncsrc/accelerate/accelerator.py:1229-1237)是节拍器本体:

if self.gradient_state.sync_with_dataloader and self.gradient_state.end_of_dataloader:
self.step = 0
self.gradient_state._set_sync_gradients(True) # 数据见底:强制同步收尾
else:
self.step += 1
self.gradient_state._set_sync_gradients((self.step % self.gradient_state.num_steps) == 0)

注意第一支:epoch 尾部不足一个完整累积周期时强制同步,保证梯度不丢。

5.5 no_sync:省掉中间步的 all-reduce

Accelerator.accumulatesrc/accelerate/accelerator.py:1255-1298)先调 _do_sync,再按 allow_gradient_sync 决定进不进 self.no_sync(m)no_syncsrc/accelerate/accelerator.py:1132-1178):FSDP2 用 model.set_requires_gradient_sync(False);DDP 直接用 model.no_sync;DeepSpeed ZeRO stage ≥ 2 时不退化成 no_sync(它的梯度分片语义不同,context 保持 nullcontext)。非分布式时整个 context 是 nullcontext(),单卡累积零开销。

5.6 optimizer 与 scheduler 的门控

AcceleratedOptimizer.stepsrc/accelerate/optimizer.py:145-178):

  • XLA 且梯度未同步时先手动 all-reduce(src/accelerate/optimizer.py:149-156);
  • sync_gradients=False 直接返回(空转);
  • 有 scaler 时把 optimizer.step 临时替换成带标记的补丁方法,scaler.step(self.optimizer, closure)scaler.update();补丁没被调用说明 scaler 检测到溢出跳过了真 step,记下 _is_overflowsrc/accelerate/optimizer.py:162-177)。补丁机制本身是 patch_optimizer_stepsrc/accelerate/optimizer.py:208-213)。

AcceleratedScheduler.stepsrc/accelerate/scheduler.py:54-84):

  • sync_gradients=False 时不步进;若插件开了 adjust_scheduler,只把内部 _step_count + 1 补偿(src/accelerate/scheduler.py:62-65);
  • 任一 optimizer step_was_skipped(fp16 溢出)则不步进(src/accelerate/scheduler.py:67-69)——LR 不会因为溢出空步而「走快」;
  • split_batches=False 时全局 batch 被放大 num_processes 倍,每次步进连走 num_processes 步(src/accelerate/scheduler.py:76-84),OneCycle 这类有 total_steps 的做边界保护。这一手让用户按原始 batch 写的 scheduler 配置在多卡下仍然正确。

6. 坑

  1. 自己写了 autocast 又开了 mixed_precision = 双重嵌套。模型的 forward 已被包进 autocast;评估时要关 autocast 请用 accelerator.autocast()src/accelerate/accelerator.py:4178-4199)并可传 handler 覆盖,不要裸写 torch.autocast 猜状态。
  2. 绕过 accelerator.backward 直接 loss.backward():loss 没除累积步数(梯度偏大 N 倍)、fp16 没走 scaler(梯度可能下溢)——两个开关同时失效。
  3. DeepSpeed 下语义外包。loss 缩放、梯度裁剪、累积边界都由 engine 接管,accelerator.backward 只是转调;读 Accelerate 源码理解 DeepSpeed 行为会南辕北辙。
  4. sync_each_batch=True 的权衡。它让 accumulate 不再进 no_sync(src/accelerate/accelerator.py:1281-1290),每步都 all-reduce——显存峰值降、速度降,适合显存极度紧张的累积场景。
  5. epoch 尾部的强制同步sync_with_dataloader=True(默认)会在 dataloader 见底时强制 sync_gradients=True 并清零 step 计数(src/accelerate/accelerator.py:1231-1233)——跨 epoch 的累积计数不会延续,这是特性不是 bug;要延续得显式关掉它。
  6. scheduler 的 num_processes 补步与 step_scheduler_with_optimizer=False 互斥。后者让 scheduler 与 optimizer 脱钩(每次调用都步进),多卡下就别再依赖补步语义,步数要自己算。