跳到主要内容

数据截至 (上游 commit fd01e35c83d8)

02 · 状态单例与进程启动

这一章讲什么: N 个进程跑同一份脚本,每个进程怎么知道「我是谁、世界有多大、用哪张卡」?accelerate launchnotebook_launcher 各自怎么起进程?答案是「环境变量契约 + Borg 单例」:启动器只写环境,单例只读环境,中间的粘合剂是 torch.distributed


1. 它要解决的小问题

分布式训练的第一性问题是进程发现:8 个进程各自独立启动,它们需要一个「汇合点」来交换地址、协商 rank、绑定设备。其次是一个工程问题:库的内部代码(比如被 prepare 包装出来的 optimizer)也需要知道 num_processes、device 这些信息,难道要一路透传?

2. 直觉

Accelerate 把两件事都压到最简:

  • 启动:不自己发明进程管理,直接复用 torchruntorch.distributed.run)或 PyTorch 的 start_processes/elastic_launch。启动器的全部职责是按约定写环境变量:LOCAL_RANKRANKWORLD_SIZEMASTER_ADDRMASTER_PORT
  • 感知:用 Borg 模式——所有实例共享同一个 __dict__PartialState() 在库里被 new 了几十次,每次拿到的都是同一份状态;它在首次构造时读环境变量、调 init_process_group、算出 device,之后整个进程随处可用。

这个设计的用户体感是:你从没把 state 传给任何东西,但每个包装对象都「知道」分布式配置。

3. 图示

accelerate launch train.py(每个进程)
───────────────── ─────────────────────────
拼装环境变量 Accelerator()
(ACCELERATE_*, MASTER_*, ──▶ AcceleratorState.__init__
CUDA_VISIBLE_DEVICES) │ 共享 __dict__
│ ▼
▼ PartialState.__init__
torch.distributed.run │ 读 LOCAL_RANK 等环境变量
(torchrun 拉起 N 个进程, ▼
注入 RANK/LOCAL_RANK) _prepare_backend 探测后端
│ (xla? nccl? gloo? mpi? ...)
▼ │
N × python train.py ▼
init_process_group(backend)


set_device(): device =
cuda:local_process_index

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

Borg 单例的全部机关只有三行:

class PartialState:
_shared_state = {} # 类级共享字典

def __init__(self, cpu=False, **kwargs):
self.__dict__ = self._shared_state # 实例 dict 换成共享 dict
if not self.initialized: # 只有第一个实例会进这里
self.backend, self.distributed_type = self._prepare_backend(...)
torch.distributed.init_process_group(backend=self.backend, **kwargs)
self.num_processes = torch.distributed.get_world_size()
self.process_index = torch.distributed.get_rank()
self.set_device() # device = cuda:local_rank

@property
def initialized(self):
return self._shared_state != {}

之后任何地方 state = PartialState(),读到的 state.devicestate.process_index 都是首次初始化的结果——属性赋值也写进共享 dict,天然全局可见。

5. 真实实现

5.1 Borg 的载体:SharedDict 与 __dict__ 偷换

PartialState 定义在 src/accelerate/state.py:123。共享字典 _shared_state 的类型由 SharedDict 决定(src/accelerate/state.py:91-120):普通平台就是 dict;torch_xla 平台换成 ThreadLocalSharedDict——一个 threading.local 描述符,因为 TPU v2/v3 的 PJRT 多线程模式要求每线程一份状态。

偷换发生在 src/accelerate/state.py:178self.__dict__ = self._shared_stateif not self.initialized 保证初始化逻辑只跑一次(src/accelerate/state.py:179)。

5.2 后端探测链:_prepare_backend

src/accelerate/state.py:755 按固定优先级嗅探环境,决定 (backend, distributed_type)

  1. SageMaker DP → smddp
  2. torch_xla 可用 → xla
  3. 环境里有 LOCAL_RANK(说明是被多进程启动器拉起的)→ 按硬件依次探测 MLU(cncl) / SDAA(tccl) / MUSA(mccl) / NPU(hccl) / HPU(hccl) / CUDA(nccl) / XPU(xccl) / Neuron(src/accelerate/state.py:768-795);
  4. 强制 CPU 且有分布式迹象(LOCAL_RANKPMI_SIZE/WORLD_SIZE > 1)→ mpigloosrc/accelerate/state.py:798-811);
  5. 都不中 → DistributedType.NO,单机。

探测完回到 __init__:非 CPU/XLA 且有 LOCAL_RANK 时调 torch.distributed.init_process_group(backend=self.backend, **kwargs)src/accelerate/state.py:244);FSDP + CPU offload 的组合还会把 backend 改写为 cuda:nccl,cpu:gloo 复合形式(src/accelerate/state.py:227-236)。MULTI_CPU 路径更特殊:先补 RANK/WORLD_SIZE/MASTER_PORT 环境变量,按物理核数自动设置 OMP_NUM_THREADS,再 init_process_group(src/accelerate/state.py:248-286)。

5.3 设备绑定:set_device

src/accelerate/state.py:819。单机直接用 default_device;分布式下把 distributed_type 字符串末段映射成 torch 设备名(MULTI_GPU → cuda),设备序号取 local_process_index % device_count(),并调 device_module.set_device 固定当前进程用卡(src/accelerate/state.py:837-843)。

5.4 训练态单例:AcceleratorState

src/accelerate/state.py:868。它先 self.__dict__.update(PartialState._shared_state) 把进程态并进来(src/accelerate/state.py:917),再叠加训练态:混合精度决策(fp8 算力不足时降级 fp16,Gaudi1 降级 bf16,src/accelerate/state.py:929-947)、DeepSpeed/FSDP/Megatron 插件的激活与环境变量回写(ACCELERATE_USE_DEEPSPEED 等,src/accelerate/state.py:973-1020)、dynamo 非混合精度时打开 TF32(src/accelerate/state.py:1028-1033)。

注意它有个防呆:不是从 Accelerator 构造(_from_accelerator=False)会直接 raise(src/accelerate/state.py:950-954),因为混精度和插件只能由 Accelerator 定夺。

第三个单例 GradientStatesrc/accelerate/state.py:1231)管梯度累积节拍,见第 4 章。

5.5 启动路径一:accelerate launch(CLI)

CLI 入口在 src/accelerate/commands/launch.py,按参数分发到若干 launcher:

  • 单机单进程 simple_launchersrc/accelerate/commands/launch.py:986):拼好环境后 subprocess.Popen 起子进程。
  • 多 GPU multi_gpu_launchersrc/accelerate/commands/launch.py:998):调 prepare_multi_gpu_envsrc/accelerate/utils/launch.py:201)算端口(占用时单机自动改 standalone 模式找空端口,src/accelerate/utils/launch.py:236-245)、写 CUDA_VISIBLE_DEVICESACCELERATE_MIXED_PRECISION 等;然后把参数过滤成 torchrun 的参数,交给 torch.distributed.runsrc/accelerate/commands/launch.py:1020-1023)。所以 accelerate launch 本质是 torchrun 的预处理壳,rank 注入由 torchrun 完成。
  • DeepSpeed deepspeed_launchersrc/accelerate/commands/launch.py:1033):多机时写 DeepSpeed 的环境文件后用其 launcher,单机仍走 torchrun。
  • TPU tpu_launchersrc/accelerate/commands/launch.py:1086):走 torch_xlaxmp.spawn

5.6 启动路径二:notebook_launcher(进程内 fork)

src/accelerate/launchers.py:43。用于 Colab/Kaggle,没有外部启动器,只能自己起进程:

  • TPU:xmp.spawn(launcher, start_method="fork")src/accelerate/launchers.py:153);
  • 多设备:用 patch_environment 写入 nproc/node_rank/world_size/master_addr/master_port/mixed_precision(src/accelerate/launchers.py:191-198),再构造 LaunchConfigelastic_launchsrc/accelerate/launchers.py:261);XPU/ROCm 必须用 spawn,其余用 fork(src/accelerate/launchers.py:213)。

无论哪条路,子进程入口都是 PrepareForLaunchsrc/accelerate/utils/launch.py:783)——它的 __call__(index)LOCAL_RANK=indexRANK=nproc*node_rank+index 写进环境(src/accelerate/utils/launch.py:821-824),再调用户的训练函数。至此与 CLI 路径合流:子进程里的 PartialState 靠同一套环境变量完成初始化。

另有 debug_launchersrc/accelerate/launchers.py:287):CPU + gloo + FileStore 起多进程,是库的测试基础设施,也是本地复现分布式 bug 的最快方式。

6. 坑

  1. notebook 里提前碰 CUDA,fork 必炸。CUDA context 不能跨 fork 继承,notebook_launcher 会检查 AcceleratorState 是否已初始化、bitsandbytes 这类「import 即初始化设备」的库是否已加载,命中即报错(src/accelerate/launchers.py:174-186)。ACCELERATE_DEBUG_MODE=1 可以先跑一次 dummy launch 预检。
  2. RTX 4000 系老驱动的 P2P/IB 问题PartialState 初始化时若检测到且未设置 NCCL_P2P_DISABLE/NCCL_IB_DISABLE,直接 raise 让你回去设环境变量(src/accelerate/state.py:318-326);CLI 路径则会自动补上。
  3. _reset_state() 之后的幽灵报错。测试重置单例后再访问旧引用,会命中 __getattr__ 的定制报错(src/accelerate/state.py:855-865)——报错文案是提示你重新初始化,不是在说属性不存在。
  4. MULTI_CPU 的隐式环境改写。它会回头写 RANK/WORLD_SIZE 等环境变量并自动设置 OMP_NUM_THREADSsrc/accelerate/state.py:248-278);多机时若没设 MASTER_ADDR 会拿到一个「请导出 rank 0 主机名」的报错。
  5. fork 启动的进程不 destroy 进程组destroy_process_groupfork_launched 直接跳过(src/accelerate/state.py:845-853),因为进程组由父进程/elastic agent 管理;自己写清理逻辑时注意别重复销毁。
  6. 设备序号取模的边界local_process_index % device_count()src/accelerate/state.py:841)意味着单机进程数超过卡数时会静默复用卡——不会报错,但显存会爆。