第 8 讲:并行基础(系统细节)

目录 · ← l7 · l9 →

第 8 讲:并行基础(系统细节)

日期:4 月 22 日(周三,Spring 2026) | 讲师:Tatsu Hashimoto | 材料:lecture_08.pdf(第 7 讲的并列 PDF 位于私有仓库,标注为受限资源)

概览

第 7 讲搭出了最简分布式代码,本讲给出系统层深挖:网络基础(为什么不能把所有设备全连起来)、并行策略全景——朴素 DDP、ZeRO 第 1–3 级(FSDP)、流水线、张量、序列/上下文、专家并行——并为每种策略做显存与通信量核算,最后总结真实大规模训练(DeepSeek、Llama 3 405B、Gemma 2、Mixtral、Qwen 3、Nemotron)如何组合它们(”3D/4D 并行”)。

核心概念与定义

  • 为什么要多 GPU:单卡扩展同时受算力(世界最快超算也不过 exaflops 级)与显存(大模型装不下)限制。并行把显存算力一起摊到多卡/多机。节点内用高速互联,节点间走网络。
  • 网络拓扑:TPU 用环形网格(toroidal mesh,便宜,非常适合张量并行);GPU 用 all-to-all/树形拓扑(更适合非结构化通信,如专家并行)。TPU8i/8t 正向树形/交换网络演进(为 MoE)。为什么不把所有东西都连起来? 成本——域大小与物理限制。
  • 朴素 DDP 的显存核算:每个参数约需 16 字节:2(bf16 参数)+ 2(bf16 梯度)+ 4(fp32 主权重)+ 4 + 4(Adam 一阶/二阶矩)——即”权重的 5 份副本”问题。这就是 DDP 显存不随设备数扩展的原因。
  • ZeRO(Zero Redundancy Optimizer)各级——把冗余副本分片:
    • 第 1 级——优化器状态分片:把一阶/二阶矩分散到各 GPU。步骤:算出完整本地梯度 → reduce-scatter 梯度(每个 rank 拿到自己负责的切片)→ 只更新自己那片参数 → all-gather 更新后的参数。通信 = 2×#params(与 DDP 的 all-reduce 相同!),显存 = (4 + K/N_gpu)×#params。“ZeRO 第 1 级是免费的(在带宽受限区间)——显存收益白拿。”
    • 第 2 级——再加梯度分片:梯度也分片;反向过程中一旦某个梯度被归约完就立即释放(从不物化完整梯度向量)。
    • 第 3 级 = FSDP——连参数都分片:参数在前向/反向时按需 all-gather,用完即释放。通信 = 3×#params(DDP 的 1.5 倍),但纯 bf16 训练下每 rank 显存降到 12/8 字节每参数。关键技巧:增量式通信/计算重叠——在计算 W0 的同时去 all-gather W1、W2,把通信开销藏起来。
    • 类比:DDP 像每个图书馆都藏一整套全书的副本;ZeRO-3 像图书馆联盟——每个分馆只保留一个书架,谁需要就把书调来(all-gather)、用完就还回去(释放)。
  • 模型并行(分片参数,通信激活——而 ZeRO-3 通信的是参数):
    • 流水线并行(逐层):朴素的逐层并行利用率极差(每张 GPU 只有 1/n 时间在工作)。micro-batch 解决它;气泡 ≈ (n_stages−1)/n_micro。通信特性好(点对点、激活大小),用于节点间;性能高度依赖 batch 大小。”Zero-bubble”变体把反向拆成”激活反向传播”与”权重梯度计算”两部分。
    • 张量并行:沿宽度切分矩阵乘。前向:f = 恒等,g = all-reduce(把部分和相加);反向:f = all-reduce,g = 恒等。QKV/up-projection 做列切分;attention 输出/down-projection 做行切分;norm/router 复制。通信:每层 8·b·s·h·(n−1)/n(all-reduce),对比流水线的 b·s·h 点对点。用在互联快的地方(节点内,≤8 卡)。优点:没有气泡、复杂度低、不需要大 batch。
    • 序列/上下文并行:把序列维度切开,用于逐点运算(LayerNorm、dropout)与长上下文(ring attention),让激活显存随设备数下降。
    • 专家并行:把专家切到不同设备,用 all-to-all 路由 token(仅用于 MoE 的 MLP 部分)。需要每个专家有足够多的 token 才高效。
  • 激活显存:隐藏的主导项——即使参数分片做得很完美,激活仍可能压垮显存(例如包含 dropout 的二次 attention 项约 5·a·s·h,可通过重计算消除;LayerNorm/dropout 的 10·s·b·h 项可通过序列并行消除)。
  • 组合策略——”3D 并行”经验法则:(1) 在模型装得下之前:节点内做张量/专家并行,跨机器做流水线(或按带宽情况用 ZeRO-3);(2) 剩下的一路用数据并行扩;若 batch 太小,就用梯度累积,以更大 batch 换更高通信效率。示例(Narayanan 等 2021):先 TP=8,再用 PP 让模型装下,DP 随规模增大而缩小(DP:32→32→32→24→15→9→6)。
  • 真实配方:DeepSeek v3:ZeRO-1 + TP=1 + EP=64 + PP=16;Llama 3 405B:DP=128、TP=8、PP=16;Gemma 2:ZeRO-3 + MP(TP+SP) + DP=768;Mixtral 8x22B(Megatron):TP/PP/CP/EP = 4/4/1/8;Nemotron 3 120B:TP=2、EP=64、CP=64;Qwen 3:EP=32、TP=2、PP=8。

代码示例:ZeRO-1(优化器状态分片)概念实现

代码(Python):

# 每个训练步,world_size 台设备,每台只负责 params[my_slice]:
# 第 1 步:每台设备在本地 batch 上算出完整梯度
loss.backward()                       # 每 rank 的 param.grad 都是完整大小

# 第 2 步:reduce-scatter 梯度 -> 每个 rank 只保留自己那一份切片
grad_slices = [torch.empty_like(grad_chunk) for ...]
dist.reduce_scatter_tensor(output=my_grad_slice, input=full_grad,
                           op=dist.ReduceOp.SUM)

# 第 3 步:每台设备只用自己那片的梯度 + 状态,更新自己负责的参数
for i in my_param_indices:
    state[i].m1 += ...                # fp32 矩只存在于本 rank
    state[i].m2 += ...
    params[i] -= lr * update(my_grad_slice_i, state[i])

# 第 4 步:all-gather 更新后的参数,让每台设备都有完整模型
dist.all_gather_into_tensor(output=full_params, input=my_param_slice)

代码做了什么: 勾画 ZeRO-1 的循环:完整本地梯度 → reduce-scatter → 在本地更新一部分参数 → all-gather 刚更新好的参数。

实现深挖:

  • 为什么通信量仍然是 2×#params:reduce-scatter 送 #params,all-gather 送 #params——恰好等于一次 all-reduce 的数据量。在带宽受限区间,ZeRO-1 相对 DDP 是”免费的”,同时把优化器状态显存降低 N_gpu 倍。
  • 为什么”先更新再聚合”而不是”先聚合再更新”:每个 rank 只需要自己那片的梯度与矩,因此在 all-gather 之前,更新是完全并行的。
  • 为什么这是作业 2”优化器状态分片”的基础:作业正是要求你实现 reduce-scatter 梯度、本地 AdamW 更新、all-gather 参数——算法完全相同。

与作业的联系:作业 2 任务:(4) 分布式数据并行(all-reduce);(5) 优化器状态分片(reduce-scatter + all-gather,如上);(6) FSDP(连参数都分片,按需 all-gather 并重叠)。讲义”ZeRO 第 3 级是 3×#param——1.5 倍通信开销,但不算差”的分析,正是你 FSDP 实现与实验报告要复现的内容。

代码示例:FSDP 风格的分片前向(概念)

代码(Python):

# 每个 rank 只保存每个权重 W_l 的一个分片。
# 要用第 l 层时,先聚合出完整权重,用完释放:
def fsdp_forward(x, layers, rank, world_size):
    for layer in layers:
        # 1. all-gather 本层的权重分片 -> 每个 rank 拿到完整 W
        full_W = all_gather(layer.weight_shard[rank])
        # 2. 计算(可与下一次 all-gather 重叠)
        x = x @ full_W
        x = F.gelu(x)
        # 3. 释放 full_W(只保留分片)
    return x

代码做了什么: 展示 FSDP 的”按需物化参数”:只在需要时聚合某一层的权重,算完即丢——因此显存峰值只包含”一层完整权重”,而非整个模型。

实现深挖:

  • 为什么重叠是关键:如果聚合是阻塞的,FSDP 会比 DDP 更慢;通过在当前矩阵乘执行时同时发起下一次 all-gather(增量式通信/计算),通信开销被掩盖——例如 (W1W0 + W2W0)x = y 在算 W0 时就把 W1、W2 聚好。
  • 为什么通信量是 3×#params:两次 all-gather(前向与反向的参数物化)+ 一次 reduce-scatter(梯度)。讲义指出这是 DDP 流量的 1.5 倍,但换来显存随设备数线性下降。
  • 为什么”概念上非常简单——写个 FSDP block wrapper 就行”:魔法在于把每个 module 包起来,让它的参数透明地分片/聚合;作业 2 要求的正是这样一个包裹 torch.nn.ModuleFSDP 类。

与作业的联系:作业 2 的 FSDP 任务(全分片数据并行训练,含前向/反向的 gather 与梯度 reduce-scatter,并与 DDP 对比 benchmark)就是这个设计。”装得下吗?”表格(在 8×A100-80G、每参数 12 字节下:基线 6.67B → ZeRO-1 16B → ZeRO-2 24.6B → ZeRO-3 53.3B)就是你在报告中要复现的显存算术。

关键要点

  1. 朴素 DDP 显存效率低(约 16 字节/参数);ZeRO 第 1→2→3 级依次分片优化器状态、梯度、参数——第 3 级只多 1.5 倍通信,却换来线性显存扩展。
  2. ZeRO-1 是”免费的”:通信量与 DDP 相同、显存严格更优——”所以你不如总是开着它”。
  3. 模型并行是”分片参数、通信激活”:流水线(深度、点对点、气泡)、张量(宽度、每层 all-reduce、需 NVLink)、序列(长度)、专家(MoE 路由、all-to-all)。
  4. 真实训练把上述手段全部组合:节点内 TP ≤ 8、跨节点 PP、其余用 DP、MoE 层用 EP、长上下文用 CP——并且处处做通信/计算重叠。
  5. 显存是动态的:激活常常超过参数;重计算与序列并行是主要杠杆。

常见陷阱

  • 显存受限却仍用 DDP:每 rank 复制整个模型;应改用 ZeRO-1/2(几乎免费)或 FSDP(1.5 倍通信)。
  • 不做通信与计算重叠:无重叠的 FSDP 在墙钟时间上严格劣于 DDP;它的全部意义就是把 gather 延迟藏起来。
  • 阻塞式流水线发送:没有 1F1B 式调度与异步操作,流水线会停摆。
  • 在慢链路上做 TP:跨以太网的逐层 all-reduce 会摧毁吞吐;TP 应限制在节点内。
  • 忽略激活显存:参数分片做到完美仍可能因为激活而 OOM;要用序列并行 + 重计算。
  • 天真地组合 DP 与 EP:DP 通常与 EP 共享副本(因此 EP < DP),而 DP 与 TP 组合不当会降低利用率。

复习题

  1. 问: 为什么说 ZeRO-1 相对 DDP 是”免费的”?
    • 答: 它的通信量是 2×#params——与 DDP 那一次 all-reduce 相同——因为 reduce-scatter + all-gather 的总搬运字节数与 all-reduce 相同。但显存从 (4+K)×#params 降到 (4+K/N_gpu)×#params。带宽代价相同,显存严格更优。
  2. 问: FSDP 聚集什么、什么时候聚集?为什么 3×#params”不算差”?
    • 答: FSDP 在前向和反向中按层按需 all-gather 参数分片(2×#params),并对梯度做 reduce-scatter(1×#params):合计 3×#params,是 DDP 的 2×#params 的 1.5 倍——但它把所有显存都分片了(参数、梯度、优化器状态,配合序列并行还包括激活)。
  3. 问: 为什么张量并行比流水线并行需要更快的互联?
    • 答: TP 每层都要通信(激活大小的张量做 all-reduce,每层 8·b·s·h·(n−1)/n);流水线只在 stage 之间通信(每 micro-batch 一次 b·s·h 的点对点)。TP 的逐层、近似全连接式流量需要 NVLink;流水线的稀疏点对点可以跑在 Infiniband 上。