Back to writing

/ CS336

[CS336-08] Distributed Training in Code

Notes for Stanford CS336 Spring 2025 lecture 8: torch.distributed, NCCL.

1 minCS336 · torch.distributed · NCCL

这篇是 Stanford CS336 Spring 2025 第 8 讲 Course Materials 的学习笔记。它把上一讲的并行概念落到 PyTorch distributed 代码。

1. 从概念到 collective API#

第 8 讲的前半部分是分布式通信积木:broadcast、scatter、gather、reduce、all-gather、reduce-scatter、all-reduce。

这些 API 看起来简单,但背后对应训练系统里的不同数据流:

  • broadcast:把参数或配置从一个 rank 分发出去;
  • all-reduce:同步 data parallel gradients;
  • reduce-scatter:同步并切分梯度;
  • all-gather:收集分片参数或 activation;
  • scatter/gather:显式分发和回收不同 shard。

掌握这些 collective,比直接背 DDP/FSDP 更底层,因为很多高级库最终都是组合这些原语。

2. NCCL 和 torch.distributed#

在 NVIDIA GPU 集群里,NCCL 是 collective communication 的关键后端。PyTorch 的 torch.distributed 提供统一接口,底层可以调用 NCCL 来执行 GPU 间通信。

分布式程序通常需要初始化 process group,给每个进程分配 rank 和 world size,再让不同 rank 参与 collective。

这带来一个重要 mental model:分布式训练不是一个 Python 进程控制所有 GPU,而是多个进程运行相同或相似代码,通过 rank 分工和 collective 协作。

3. All-reduce 为什么等于训练同步#

Data parallel 中,每个 rank 拿到不同 batch shard,计算本地 gradient。为了让所有模型副本保持一致,需要把各 rank gradient 求和或平均。

这就是 all-reduce 的作用:

每个 rank 输入自己的 gradient
每个 rank 输出全局平均 gradient

然后每个 rank 用相同 gradient 更新自己的模型参数。这样虽然数据不同,但参数保持一致。

4. reduce-scatter 和 all-gather 的意义#

ZeRO 和 optimizer state sharding 的关键是不要让每张卡都保存所有状态。

reduce-scatter 可以在做 gradient reduce 的同时,把结果切成 shard 分给不同 rank。每个 rank 只负责一部分参数的优化器更新。更新后,再用 all-gather 让需要完整参数的阶段拿到对应数据。

这解释了为什么 reduce-scatter/all-gather 是现代分布式训练的核心积木。

5. 分布式训练的工程风险#

写 distributed code 时,错误往往不明显:一个 rank 进了 collective,另一个 rank 没进,就会 hang;shape 或 dtype 不匹配,可能报错也可能卡住;通信和计算不同步,会导致 profiler 难读。

因此实践中需要特别关注:

  • 所有 rank 是否执行相同 collective;
  • tensor shape、dtype、device 是否一致;
  • 是否正确设置 environment variables;
  • 是否在 benchmark 中同步 CUDA;
  • 通信是否和计算重叠。

6. takeaway#

第 8 讲的价值是把“并行训练”拆成可执行的通信原语。

  • DDP 的核心是 gradient all-reduce;
  • ZeRO 的核心是 reduce-scatter 和 all-gather;
  • NCCL 是 GPU collective 的高性能后端;
  • 分布式程序的基本单位是 rank,而不是单机脚本。

理解这些之后,再看 Megatron、DeepSpeed、FSDP、vLLM 的分布式路径,会更容易定位通信发生在哪里。

参考#