数据截至 (上游 commit fd01e35c83d8)
Accelerate — 架构与原理
30 秒导读: Accelerate 是 HuggingFace 的「分布式训练与大模型推理适配层」。它不发明任何训练算法,只做两件事:第一,让你用同一份纯 PyTorch 训练脚本跑在单卡、多卡 DDP、FSDP、DeepSpeed、TPU 上——
accelerator.prepare(model, optimizer, dataloader, scheduler)一行把四样对象换成带分布式和混合精度能力的包装版;第二,让你把塞不进单卡的大模型按层切到多张 GPU、CPU 内存甚至磁盘上跑推理(device_map="auto")。整个库的主干不到二十个文件,核心技巧是「包装 + 钩子」:不改你的类,只在对象外面 套一层。
1. 这是什么(零基础也能懂)
一句话定义
Accelerate 是一个 PyTorch 训练/推理代码的「运行环境适配层」:你写单卡的 PyTorch 代码,它负责把这份代码在 N 个进程、M 种硬件、K 种分布式后端下跑起来,行为保持一致。
它要解决谁的什么问题
假设你写好了一个单卡训练循环,现在要上 8 卡:
- 得自己起 8 个进程、配
init_process_group、给每个进程绑卡; - 得把模型包成
DistributedDataParallel,把 DataLoader 换成DistributedSampler版本,还要处理「数据集不能整除进程数」的尾巴; - 上混合精度要包
autocast+GradScaler,梯度累积要处理 DDP 的no_sync、scaler 溢出时跳过 step; - 想换 FSDP 或 DeepSpeed?以上大半要重写。
Accelerate 把这些「环境差异」全部吸收到 prepare() 和三个状态单例里。你的训练循环只调 accelerator.backward(loss)、accelerator.gather(x) 这类中性 API。
另一条独立的产品线是大模型推理:70B 模型放不进一张 24GB 的卡,Accelerate 提供 init_empty_weights() + device_map="auto" + dispatch_model(),把层切成「GPU 常驻 + CPU/磁盘暂存、用到才搬上来」的结构。HuggingFace transformers 的 device_map 参数底层就是这套机制。
三个边界声明
- 不是训练框架:没有 Trainer、没有训练循环抽象,循环永远是你自己的(transformers 的 Trainer 才是训练框架,其内部正是调用 Accelerate)。
- 不是分布式引擎:真正的 all-reduce、分片、ZeRO 都由 PyTorch 分布式、FSDP、DeepSpeed、Megatron-LM 完成,Accelerate 只负责「按配置把对象交给正确的引擎」。
- device_map 是推理机制:用它加载的模型不能直接分布式训练(
prepare会显式拒绝,见第 1、5 章)。
2. 顶层全景
2.1 一张图看全貌
accelerate launch ──▶ PartialState / AcceleratorState ──▶ Accelerator.prepare(*args)
(torchrun 起 N 进程 (Borg 单例:读环境变量、 │
写 LOCAL_RANK 等) init_process_group、定 device) │ 按类型两趟分发
┌──────────────┬─────────────────────┼──────────────────┐
▼ ▼ ▼ ▼
prepare_model prepare_optimizer prepare_data_loader prepare_scheduler
(AMP 改写 forward (Accelerated- (BatchSamplerShard / (Accelerated-
+ DDP/FSDP 包壳) Optimizer) DataLoaderDispatcher) Scheduler)
推理支线(与训练解耦的另一套 API):
init_empty_weights ──▶ infer_auto_device_map ──▶ load_checkpoint_in_model ──▶ dispatch_model ──▶ 每次 forward 经 AlignDevicesHook
(参数全在 meta 设备, (按显存预算装箱, (逐 shard 加载,按 map (给每个 block 挂 pre/ (pre: 物化权重到执行设备
零内存占用) 决定每层去哪台设备) 直接落到目标设备) post forward hook) post: 卸回 meta)
2.2 部件表
| 部件 | 位置 | 职责 |
|---|---|---|
Accelerator | src/accelerate/accelerator.py | 用户主入口:prepare / backward / gather / accumulate / save_state |
PartialState | src/accelerate/state.py:123 | Borg 单例:进程组、rank、world size、device、分布式类型 |
AcceleratorState | src/accelerate/state.py:868 | 在 PartialState 之上加混合精度与后端插件状态 |
GradientState | src/accelerate/state.py:1231 | 梯度累积节拍与 dataloader 尾部状态的单例 |
AcceleratedOptimizer | src/accelerate/optimizer.py:38 | 包装 optimizer :GradScaler step、累积期跳过、XLA 梯度同步 |
AcceleratedScheduler | src/accelerate/scheduler.py:25 | 包装 scheduler:只在 optimizer 真正 step 时步进 |
BatchSamplerShard / DataLoaderShard / DataLoaderDispatcher | src/accelerate/data_loader.py:110 / :510 / :723 | 数据分片的两种策略与迭代器 |
multi_gpu_launcher / notebook_launcher / PrepareForLaunch | src/accelerate/commands/launch.py:998、src/accelerate/launchers.py:40、src/accelerate/utils/launch.py:783 | 三条进程启动路径 |
init_empty_weights / dispatch_model / load_checkpoint_and_dispatch | src/accelerate/big_modeling.py:62 / :315 / :520 | 大模型推理三件套 |
AlignDevicesHook / add_hook_to_module | src/accelerate/hooks.py:242 / :147 | forward 前后的权重物化/卸载钩子及其注入机制 |
infer_auto_device_map / get_balanced_memory | src/accelerate/utils/modeling.py:1304 / :931 | device_map 的自动装箱算法与显存预算 |
gather / reduce / broadcast / send_to_device | src/accelerate/utils/operations.py | 跨后端集合通信与嵌套结构张量搬运原语 |
| 插件 dataclass 群 | src/accelerate/utils/dataclasses.py | DDP/FSDP/DeepSpeed/Megatron/梯度累积/FP8 等配置对象 |
2.3 主线:一次多卡训练是怎么跑 起来的
- 启动:
accelerate launch --multi_gpu --num_processes 8 train.py拼好环境变量后转交torch.distributed.run(torchrun),拉起 8 个进程(src/accelerate/commands/launch.py:998)。 - 感知:每个进程执行到
Accelerator()时,PartialState读LOCAL_RANK等环境变量,探测后端(nccl/gloo/xla/…),调用init_process_group,算出本进程的 device(src/accelerate/state.py:177)。 - 包装:
accelerator.prepare(model, optimizer, dataloader, scheduler)按类型分两趟分发:model 被改写 forward(混合精度)并按需包 DDP/FSDP;optimizer/scheduler 各套一层包装;dataloader 换成分片版(src/accelerate/accelerator.py:1414)。 - 循环:训练循环不变,只是
loss.backward()换成accelerator.backward(loss)(内部处理 loss 缩放与 scaler);评估用accelerator.gather()聚合各进程结果。 - 收尾:
accelerator.wait_for_everyone()对齐进程,unwrap_model()剥掉包装后save_state()存档。
推理支线则是:空壳初始化 → 自动装箱出 device_map → 逐 shard 加载到目标设备 → 给每个 block 挂钩子,forward 时逐层「物化→计算→卸载」。