跳到主要内容

数据截至 (上游 commit fd01e35c83d8)

01 · prepare 包装机制

这一章讲什么: accelerator.prepare(model, optimizer, dataloader, scheduler) 这一行内部分两趟做了什么,四类对象各自被换成了什么,以及为什么这套包装对训练循环是透明的。DataLoader 分片的细节单独立在第 3 章,本章只管调度与包装骨架。


1. 它要解决的小问题

你写了一份纯 PyTorch 训练脚本。现在要求:同一份代码,今天在单卡调试,明天在 8 卡 DDP 上跑,下周可能切 FSDP 或 DeepSpeed——而训练循环一行不改。

要做到这一点,model、optimizer、dataloader、scheduler 这四样对象必须「看起来没变,用起来变了」:模型 forward 时自动进 autocast、反向时梯度自动 all-reduce;optimizer.step 在梯度累积的中间步自动空转;dataloader 每个进程只出自己的那份数据。prepare() 就是完成这次「偷梁换柱」的地方。

2. 直觉

Accelerate 的选择是拦截点包装:不发明 AcceleratedModel 这种新类型要求你继承,而是在你已有对象的外面套一层壳,壳上实现分布式语义,壳的接口与原对象一致。

这带来两个推论:

  • 包装必须按类型分发——四种对象的「壳」完全不同,prepare 要先认出每个参数是什么。
  • 包装有先后顺序——scheduler 依赖 optimizer(它的 step 要看 optimizer 是否真的走了步),所以 optimizer 必须先包完。

3. 图示

prepare() 的两趟调度(以通用路径为例):

prepare(model, optimizer, dataloader, scheduler)

│ 第一趟 _prepare_one(first_pass=True)
├──────────┬──────────────┬───────────────┬──────────────┐
▼ ▼ ▼ ▼ ▼
prepare_ prepare_ prepare_ (scheduler 不认识的对象
model optimizer data_loader 跳过) 原样返回
│ │ │
▼ ▼ ▼
AMP 改写 Accelerated- BatchSamplerShard /
forward Optimizer DataLoaderDispatcher
+ DDP/FSDP

│ 第二趟 _prepare_one(first_pass=False)

prepare_scheduler(此时已能从 self._optimizers 里找到包装后的 optimizer)


全部打上 _is_accelerate_prepared 标记,按原顺序返回

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

prepare 的核心逻辑可以压缩成这样一段伪代码(省略了所有后端特判):

class Accelerator:
def _prepare_one(self, obj, first_pass=False):
if first_pass: # 第一趟:model / optimizer / dataloader
if isinstance(obj, DataLoader):
return self.prepare_data_loader(obj)
if isinstance(obj, torch.nn.Module):
return self.prepare_model(obj)
if isinstance(obj, torch.optim.Optimizer):
return self.prepare_optimizer(obj)
elif isinstance(obj, LRScheduler): # 第二趟:scheduler
return self.prepare_scheduler(obj)
return obj # 不认识的对象原样穿过

def prepare(self, *args):
result = tuple(self._prepare_one(o, first_pass=True) for o in args)
result = tuple(self._prepare_one(o) for o in result)
for item in result:
item._is_accelerate_prepared = True # 幂等标记
return result

def prepare_model(self, model):
if self.native_amp: # 混合精度:改写 forward
model._original_forward = model.forward
model.forward = convert_outputs_to_fp32(autocast_ctx(model.forward))
model = model.to(self.device) # 设备放置
if self.multi_device: # 分布式:包 DDP
model = DistributedDataParallel(model, device_ids=[self.local_process_index])
return model

要点只有两个:按类型分两趟分发;每种 prepare_* 只做「设备放置 + 套壳」。

5. 真实实现

5.1 入口与校验:Accelerator.prepare

入口在 src/accelerate/accelerator.py:1414。在分发之前,它先做一轮「红线检查」:

  • hf_device_map 的模型禁止在分布式模式下训练,直接 raise(src/accelerate/accelerator.py:1469-1480);
  • DeepSpeed 模式同一个 Accelerator 只允许一个模型(src/accelerate/accelerator.py:1483-1492);
  • XLA 上模型与 optimizer 参数设备必须一致,否则说明用户在 model.to(device) 之后才建 optimizer,参数已脱节(src/accelerate/accelerator.py:1500);
  • FSDP2 要求 model 与 optimizer 必须一起传入 prepare(optimizer 要在模型转换后重建参数引用),且只支持单模型(src/accelerate/accelerator.py:1516-1523)。

之后按后端分派:TP/CP 并行先走 _prepare_tp/_prepare_cp,FP8 走 TE/AO 的换算,DeepSpeed/Megatron/FSDP2 各有专属方法,其余全部走通用的两趟 _prepare_onesrc/accelerate/accelerator.py:1557-1560)。

5.2 两趟调度:_prepare_one

src/accelerate/accelerator.py:1397 就是第 4 节伪代码的真实版,逐行对应。两趟调用在 src/accelerate/accelerator.py:1559-1560:第一趟 first_pass=True,第二趟默认。

收尾时,凡是被收进 self._dataloaders / _models / _optimizers / _schedulers 的对象,统一打上 item._is_accelerate_prepared = Truesrc/accelerate/accelerator.py:1571-1577),最后 return result if len(result) > 1 else result[0]——传一个对象就返回单个,不传 tuple。

5.3 模型:prepare_model

src/accelerate/accelerator.py:1769。步骤顺序:

  1. 幂等短路:已打过标记的模型直接登记并返回(src/accelerate/accelerator.py:1799-1803)。
  2. AMP 改写 forwardnative_amp 为真时,把原 forward 存到 model._original_forward,再把 model.forward 换成 convert_outputs_to_fp32(autocast_context(model.forward))src/accelerate/accelerator.py:1818-1827)。对 bound method 没有 __func__ 的特例用 MethodType 处理。注意这是实例级改写,类本身没动。
  3. 量化模型红线:8bit/4bit + 多设备的组合直接报错(src/accelerate/accelerator.py:1835-1841)。
  4. 设备放置device_placement 为真且无 device_map 时 model = model.to(self.device)src/accelerate/accelerator.py:1875)。
  5. 分布式包壳:多设备且含可训练参数时包 DistributedDataParalleldevice_ids/output_devicelocal_process_indexddp_handler 的 kwargs 与 comm hook 在此注入(src/accelerate/accelerator.py:1883-1896)。
  6. FSDP 分支:从 fsdp_plugin 组装 kwargs(auto_wrap_policy、mixed_precision_policy、param_init_fn 等),model = FSDP(model, **kwargs),可选激活 checkpoint 包装(src/accelerate/accelerator.py:1912-1979)。

5.4 优化器:prepare_optimizerAcceleratedOptimizer

src/accelerate/accelerator.py:2733 把 optimizer 包成 AcceleratedOptimizersrc/accelerate/optimizer.py:38)。包装层做三件事:

  • 构造时把 optimizer 的 state_dict 搬到目标设备再 load 回去(动量等状态上卡),move_to_device 递归处理嵌套 dict/list/tensor(src/accelerate/optimizer.py:30-36src/accelerate/optimizer.py:70-76);
  • step() 里按 gradient_state.sync_gradients 决定真走还是空转,有 scaler 时改走 scaler.step/update 并检测溢出(细节见第 4 章);
  • state/param_groups/defaults 等属性全部透传(src/accelerate/optimizer.py:78-103),所以包装后 optimizer.param_groups[0]["lr"] = ... 这类写法照常工作。

5.5 调度器:prepare_schedulerAcceleratedScheduler

src/accelerate/accelerator.py:2777。它先在 self._optimizers 里反查 scheduler 绑定的那个 optimizer(getattr(scheduler, "optimizer", None) == opt.optimizersrc/accelerate/accelerator.py:2803-2808),再包成 AcceleratedSchedulersrc/accelerate/scheduler.py:25)——这就是 scheduler 必须放在第二趟的原因。其 step 门控逻辑见第 4 章。

5.6 拆包装:unwrap_modelextract_model_from_parallel

存档前要剥壳:Accelerator.unwrap_modelsrc/accelerate/accelerator.py:3213)转调 extract_model_from_parallelsrc/accelerate/utils/other.py:248)。后者依次处理:compiled 模型取 _orig_mod;循环剥掉 DDP/DataParallel/DeepSpeedEngine/FSDP 的 .modulekeep_fp32_wrapper=False 时沿 __wrapped__ 链找回 _original_forward 并恢复(src/accelerate/utils/other.py:313-321)。

6. 坑

  1. 二次 prepare 是静默短路find_batch_size 等内部机制依赖这一点重复 prepare,但如果你以为「再 prepare 一次能换新配置」,什么都不会发生——对象已带 _is_accelerate_prepared,直接原样返回。
  2. device_map 模型 + 分布式 = 显式 raise。想用 device_map="auto" 加载再大训,门在 prepare 入口就关了;唯一出路是单进程跑(naive pipeline),或换 FSDP/DeepSpeed 路线。可用环境变量 ACCELERATE_BYPASS_DEVICE_MAP=true 绕过,但后果自负。
  3. XLA 的 optimizer 顺序陷阱。TPU 上 model.to(device) 会创建新参数对象,先建 optimizer 再 to,optimizer 就握着一堆「死参数」。prepare 会检测并报错,但修法是把 model.to 删掉交给 prepare,或保证顺序。
  4. FSDP2 必须 model+optimizer 同传。分开 prepare 会在入口被拒(src/accelerate/accelerator.py:1516-1523);这是 FSDP2 改参数引用(DTensor 化)的硬约束。
  5. 包装层改写的是实例属性model.forward 被换、_is_accelerate_prepared 被打在对象上——深拷贝、pickle、torch.save(model) 整模型序列化时都可能带上这些私货;存档请走 unwrap_model + save_state
  6. 多模型场景的限制。DeepSpeed 下同一 Accelerator 多模型直接 AssertionError;通用路径虽允许多模型,但 accumulate(*models)clip_grad_norm_ 等 API 的语义都要按「传进来的那批模型」理解。