分布式训练——如何把大模型分到多张卡上 / Distributing Large-Model Training Across Multiple GPUs
📅 创建时间:2026-06-02 🏷️ 标签:#分布式 #数据并行 #模型并行 #张量并行 #流水线并行 #FSDP #NCCL 📚 前置知识:[[01-gpu-hardware]](GPU 硬件基础) 📚 相关知识:[[03-memory-optimization]](显存优化) [[04-mixed-precision]](混合精度)
场景:单卡跑不了 70B 模型,必须分布式
┌─────────────────────────────────────────────────────────────┐
│ │
│ 你的模型:Llama-3-70B。 │
│ │
│ 单卡 A100 80GB 能跑吗? │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 模型参数:70B × 2 bytes(BF16)= 140 GB │ │
│ │ 梯度:70B × 2 bytes = 140 GB │ │
│ │ 优化器状态:70B × 4 bytes(FP32)= 280 GB │ │
│ │ Activation:~100 GB+(取决于 seq length 和 batch) │ │
│ │ │ │
│ │ 总计:660 GB+ → 单卡 80GB 远远不够 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 方案 1:买更大的卡? │
│ → H100 80GB 也装不下,H200 141GB 也勉强 │
│ → TB 级显存只有 HBM 集群 │
│ │
│ 方案 2:分布式训练 │
│ → 把模型切分到多张卡上 │
│ → 这是工业界唯一可行的方案 │
│ │
│ 但怎么切?每种切法有什么 trade-off? │
│ 这就是本章要回答的问题。 │
│ │
└─────────────────────────────────────────────────────────────┘第1节:数据并行(Data Parallelism)——最简单的并行方式
核心思想:每个卡都有完整的模型
┌─────────────────────────────────────────────────────────────┐
│ 数据并行:每张卡一个完整模型副本 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 假设:8 张 GPU,1 个模型 │
│ │
│ GPU 0 GPU 1 GPU 2 GPU 3 GPU 4 GPU 5 GPU 6 GPU 7 │
│ ┌───┐ ┌───┐ ┌───┐ ┌───┐ ┌───┐ ┌───┐ ┌───┐ ┌───┐ │
│ │ M │ │ M │ │ M │ │ M │ │ M │ │ M │ │ M │ │ M │ │
│ │== │ │== │ │== │ │== │ │== │ │== │ │== │ │== │ │
│ │70B│ │70B│ │70B│ │70B│ │70B│ │70B│ │70B│ │70B│ │
│ └───┘ └───┘ └───┘ └───┘ └───┘ └───┘ └───┘ └───┘ │
│ 70B 70B 70B 70B 70B 70B 70B 70B │
│ │
│ M = 完整模型(参数 + 梯度 + 优化器),8 份副本 │
│ │
│ 问题:每张卡都要存完整的 70B 模型 │
│ → 参数 140GB + 梯度 140GB + 优化器 280GB = 560GB │
│ → 单卡 80GB × 8 = 640GB → 勉强够用 │
│ → 但 70B 的 Activation 会在前向时爆炸 │
│ │
│ 数据并行解决的问题: │
│ → 显存不够?每张卡都存一份完整模型,共享数据批次 │
│ → 计算慢?8 张卡同时处理不同的数据批次 │
│ │
└─────────────────────────────────────────────────────────────┘训练流程:每张卡独立前向后向,同步梯度
┌─────────────────────────────────────────────────────────────┐
│ Data Parallel 训练流程 │
├─────────────────────────────────────────────────────────────┤
│ │
│ Step 1:分数据 │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Batch = 64,8 张卡 │ │
│ │ GPU 0: samples[0:8] GPU 1: samples[8:16] │ │
│ │ GPU 2: samples[16:24] ... │ │
│ │ 每张卡拿到 8 个样本 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 2:并行前向 + 后向(每张卡独立计算) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU 0: loss_0 = model(data_0) → backward() │ │
│ │ GPU 1: loss_1 = model(data_1) → backward() │ │
│ │ ... │ │
│ │ 所有卡同时计算,互不干扰 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 3:梯度同步(AllReduce) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ GPU 0 ←──────── AllReduce ────────→ GPU 7 │ │
│ │ grad_0 + grad_1 + ... + grad_7 │ │
│ │ 平均后写回每张卡 │ │
│ │ │ │
│ │ AllReduce 完成后,每张卡的梯度一致 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 4:更新参数(每张卡独立) │
│ optimizer.step() # 每张卡独立执行,更新后参数一致 │
│ │
└─────────────────────────────────────────────────────────────┘PyTorch DDP——数据并行的标准实现
python
import torch
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
import torch.distributed as dist
# 初始化分布式环境
dist.init_process_group(backend="nccl")
# 模型和输入要在正确的 GPU 上
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
model = MyModel().cuda()
# 包装模型——DDP 自动处理:
# 1. 梯度 AllReduce(所有卡同步)
# 2. Buffer 同步(BatchNorm 等)
# 3. 梯度累加(Gradient Bucket,分段同步减少通信开销)
model = DDP(model, device_ids=[local_rank])
# 训练循环
for data in dataloader:
# Dataloader 需要分片,每个 rank 只拿到自己的数据
data = data.cuda()
loss = model(data).sum()
loss.backward() # DDP 自动 AllReduce 所有 GPU 的梯度
optimizer.step()
optimizer.zero_grad()DDP 的关键参数和 trade-off
┌─────────────────────────────────────────────────────────────┐
│ DDP 的关键参数与影响 │
├─────────────────────────────────────────────────────────────┤
│ │
│ find_unused_parameters: │
│ → True:自动处理不参与训练的参数(如 frozen encoder) │
│ → False(默认):跳过不参与训练的分支,开销更小 │
│ │
│ gradient_as_bucket_view: │
│ → True(默认):梯度存储为 AllReduce bucket 的视图 │
│ → 节省显存,但需要在 optimizer 之前处理好 │
│ │
│ broadcast_buffers: │
│ → True(默认):前向时同步 buffer(如 BatchNorm) │
│ → False:不同 rank 可能看到不同的 buffer 状态 │
│ │
│ ⚠️ DDP 的通信:AllReduce 在 backward 后自动执行 │
│ ⚠️ 通信和计算不能 overlap(backward 必须完成后才能 AllReduce)│
│ ⚠️ 小 batch size 时,通信占比高,GPU 利用率低 │
│ │
└─────────────────────────────────────────────────────────────┘第2节:张量并行(Tensor Parallelism)——把单层切开
核心思想:单层参数按维度切分
┌─────────────────────────────────────────────────────────────┐
│ 张量并行:把单个层的参数切分到多张卡 │
├─────────────────────────────────────────────────────────────┤
│ │
│ MLP 层:Y = GeLU(XW_1)W_2 │
│ │
│ 普通情况(单卡): │
│ X (seq, batch, hidden) × W_1 (hidden, 4*hidden) │
│ → 输出 Y (seq, batch, 4*hidden) │
│ │
│ 张量并行(2 卡,列切 W_1,行切 W_2): │
│ │
│ GPU 0: GPU 1: │
│ X × W_1[:, :2*h] X × W_1[:, 2*h:] │
│ → Y_part0 → Y_part1 │
│ ↓ ↓ │
│ Y_part0 × W_2[:2*h, :] Y_part1 × W_2[2*h:, :] │
│ → Y0 → Y1 │
│ ↓ ↓ │
│ └────────── AllReduce ───────────┘ │
│ ↓ │
│ Y = Y0 + Y1 │
│ │
│ 每张卡只存 W_1 的 1/2 和 W_2 的 1/2 │
│ 但需要通信整合每层的输出 │
│ │
└─────────────────────────────────────────────────────────────┘为什么张量并行必须用 NVLink
┌─────────────────────────────────────────────────────────────┐
│ 张量并行的通信量分析 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 每个 Transformer 层的通信: │
│ │
│ Forward: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ X × W_1:每卡独立计算 │ │
│ │ AllGather Y_part → Y:所有卡需要聚合结果 │ │
│ │ 通信量:batch × seq × hidden × 2 bytes(BF16) │ │
│ │ 例:batch=1, seq=4096, hidden=8192 │ │
│ │ = 1 × 4096 × 8192 × 2 = 64 MB │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Backward: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ ReduceScatter dL/dX:反向传播梯度 │ │
│ │ 通信量同 Forward │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 每层每步:128 MB 通信(2 卡) │
│ 32 层 × 每步 2 次通信 × 若干 micro-batch: │
│ → 总通信量极大 │
│ │
│ 为什么必须 NVLink: │
│ → PCIe 50 GB/s:64 MB / 1.3 ms │
│ → NVLink 900 GB/s:64 MB / 0.07 ms │
│ → 差距 18 倍 │
│ → PCIe 会让张量并行几乎不可用 │
│ │
│ 结论: │
│ → 张量并行只适合在同一节点内(NVLink 互联) │
│ → TP = 8 意味着需要 8 张卡在同一节点 │
│ │
└─────────────────────────────────────────────────────────────┘Megatron-LM——张量并行的工业实现
python
# Megatron-LM 的张量并行配置
# tp_size = 8 表示每层参数切分到 8 张卡
from megatron.core import parallel_state
# 初始化模型并行(张量并行 + 流水线并行)
parallel_state.initialize_model_parallel(
tensor_model_parallel_size=8, # TP = 8(每层参数分到 8 卡)
pipeline_model_parallel_size=8, # PP = 8(32 层分成 8 段)
virtual_pipeline_model_parallel_size=2, # VP = 2(每个阶段 2 个 chunk)
)
# 模型会自动按 TP 切分
model = GPTModel(num_layers=32, hidden_size=8192, num_attention_heads=64)
# Forward 时,Megatron 自动处理:
# 1. ColumnParallelLinear:列切 W_1
# 2. RowParallelLinear:行切 W_2 + AllReduce
# 3. 跨层通信自动插入第3节:流水线并行(Pipeline Parallelism)——按层切分
核心思想:模型按层切分,不同卡负责不同阶段
┌─────────────────────────────────────────────────────────────┐
│ 流水线并行:模型按层切分到不同 GPU │
├─────────────────────────────────────────────────────────────┤
│ │
│ Llama-3-70B:32 层 Transformer │
│ │
│ GPU 0: [Layer 0-3] → 输出给 GPU 1 │
│ GPU 1: [Layer 4-7] → 输出给 GPU 2 │
│ GPU 2: [Layer 8-11] → 输出给 GPU 3 │
│ GPU 3: [Layer 12-15] → 输出给 GPU 4 │
│ GPU 4: [Layer 16-19] → 输出给 GPU 5 │
│ GPU 5: [Layer 20-23] → 输出给 GPU 6 │
│ GPU 6: [Layer 24-27] → 输出给 GPU 7 │
│ GPU 7: [Layer 28-31] → 输出给 GPU 8 │
│ │
│ Forward: │
│ data → GPU0 → GPU1 → GPU2 → ... → GPU7 → loss │
│ │
│ Backward: │
│ loss → GPU7 → GPU6 → ... → GPU0 → 梯度更新 │
│ │
│ 优势:每张卡只存 1/8 的层,显存压力大幅降低 │
│ 问题:数据按顺序流过,GPU 利用率低(大量空闲等待) │
│ │
└─────────────────────────────────────────────────────────────┘朴素流水线的致命问题:流水线气泡
┌─────────────────────────────────────────────────────────────┐
│ 朴素流水线的问题:大量气泡 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 假设 4 张 GPU,micro_batch=1,按层顺序执行: │
│ │
│ 时间: T1 T2 T3 T4 T5 T6 T7 T8 T9 T10 │
│ ──────────────────────────────────────────────────────────│
│ GPU0: F0 B0 W idle idle idle idle idle idle │
│ GPU1: idle F1 B1 W idle idle idle idle idle │
│ GPU2: idle idle F2 B2 W idle idle idle idle │
│ GPU3: idle idle idle F3 B3 W idle idle idle │
│ │
│ F = Forward, B = Backward, W = Weight Update │
│ │
│ 问题: │
│ → GPU0 在 T1 完成 F0 后,等待 GPU1 完成 F1 才能 backward│
│ → GPU1 在 T2 完成 F1 后,等待 GPU2 完成 F2 才能 backward│
│ → 中间出现了大量"气泡"(idle time) │
│ → 气泡时间占比 ≈ (P-1)/P(P=GPU 数),P 越大浪费越多 │
│ │
│ 4 GPU 时:气泡占比 75% │
│ 8 GPU 时:气泡占比 87.5% │
│ → 朴素流水线完全不可用 │
│ │
└─────────────────────────────────────────────────────────────┘GPipe 和 PipeDream——减少流水线气泡
┌─────────────────────────────────────────────────────────────┐
│ 1F1B + Interleaving:消除气泡 │
├─────────────────────────────────────────────────────────────┤
│ │
│ GPipe(Gradient Pipe): │
│ → 把 batch 分成多个 micro_batch,交错执行 │
│ → T1: F0 T2: F1 T3: F2 T4: F3 ← Forward 按序 │
│ → T5: B3 T6: B2 T7: B1 T8: B0 ← Backward 倒序 │
│ → 但中间的空闲等待仍然存在 │
│ │
│ 1F1B(One-Forward-One-Backward): │
│ → 尽量让每个 GPU 同时执行 F 和 B │
│ → 前一个 micro_batch backward 时,后一个 forward │
│ → 气泡减少,但仍然存在 │
│ │
│ Interleaving(交错执行): │
│ 每个 GPU 处理多个 chunk(如每 GPU 2 个 layer block) │
│ │
│ 时间: T1 T2 T3 T4 T5 T6 T7 T8 T9 T10 │
│ ───────────────────────────────────────────────────────────│
│ GPU0: F0 F1 B0 F2 B1 idle idle idle idle │
│ GPU1: idle F0 F1 B0 F2 B1 idle idle idle │
│ GPU2: idle idle F0 F1 B0 F2 B1 idle idle │
│ GPU3: idle idle idle F0 F1 B0 F2 B1 idle │
│ │
│ 气泡大幅减少,但通信量增加(chunk 数量增加) │
│ │
│ 实际选择:PP 8 + TP 8(总计 64 GPU)是常见的配置 │
│ │
└─────────────────────────────────────────────────────────────┘第4节:FSDP——全共享数据并行
ZeRO 的三个 Stage
FSDP(Fully Sharded Data Parallel)本质上是 ZeRO(Zero Redundancy Optimizer)在 DDP 基础上的实现:
┌─────────────────────────────────────────────────────────────┐
│ ZeRO Stage 1 / 2 / 3:逐步分片显存 │
├─────────────────────────────────────────────────────────────┤
│ │
│ Stage 0(基线 DDP): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU 0 GPU 1 GPU 2 GPU 3 │ │
│ │ ┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐ │ │
│ │ │ P+G+O │ │ P+G+O │ │ P+G+O │ │ P+G+O │ ← 全量副本 │ │
│ │ └─────┘ └─────┘ └─────┘ └─────┘ │ │
│ │ P=参数, G=梯度, O=优化器状态 │ │
│ └─────────────────────────────────────────────────────┘ │
│ → 显存占用:无优化(最大) │
│ │
│ Stage 1(优化器状态分片): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU 0 GPU 1 GPU 2 GPU 3 │ │
│ │ ┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐ │ │
│ │ │P+G+O1│ │P+G+O2│ │P+G+O3│ │P+G+O4│ ← 只存 1/4 优化器 │ │
│ │ │ │ │ │ │ │ │ │ │ │
│ │ │ 参数 │ │ 参数 │ │ 参数 │ │ 参数 │ ← 参数仍有副本 │ │
│ │ │ 梯度 │ │ 梯度 │ │ 梯度 │ │ 梯度 │ ← 梯度仍有副本 │ │
│ │ └─────┘ └─────┘ └─────┘ └─────┘ │ │
│ └─────────────────────────────────────────────────────┘ │
│ → 显存节省:约 4 倍(但效果有限) │
│ │
│ Stage 2(梯度分片): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU 0 GPU 1 GPU 2 GPU 3 │ │
│ │ ┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐ │ │
│ │ │P+G1 +│ │P+G2 +│ │P+G3 +│ │P+G4 +│ │ │
│ │ │ O1 │ │ O2 │ │ O3 │ │ O4 │ ← 梯度+优化器分片 │ │
│ │ │ + │ │ + │ │ + │ │ + │ │ │
│ │ │ 参数 │ │ 参数 │ │ 参数 │ │ 参数 │ ← 参数仍有副本 │ │
│ │ └─────┘ └─────┘ └─────┘ └─────┘ │ │
│ └─────────────────────────────────────────────────────┘ │
│ → 显存节省:约 8 倍 │
│ │
│ Stage 3(参数也分片): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU 0 GPU 1 GPU 2 GPU 3 │ │
│ │ ┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐ │ │
│ │ │G1+O1 │ │G2+O2 │ │G3+O3 │ │G4+O4 │ ← 梯度+优化器分片 │ │
│ │ │ +P1 │ │ +P2 │ │ +P3 │ │ +P4 │ ← 参数也分片 │ │
│ │ └─────┘ └─────┘ └─────┘ └─────┘ │ │
│ │ Forward 时 AllGather 参数,Backward 后 ReduceScatter │ │
│ └─────────────────────────────────────────────────────┘ │
│ → 显存节省:约 8 倍(每卡只存 1/N 的全部) │
│ → FSDP = ZeRO Stage 3 │
│ │
└─────────────────────────────────────────────────────────────┘FSDP vs DDP vs 张量/流水线并行对比
┌─────────────────────────────────────────────────────────────┐
│ 四种并行策略对比全景图 │
├─────────────────────────────────────────────────────────────┤
│ │
│ │ 维度 │ DDP │ 张量并行 │ 流水线并行 │ FSDP/ZeRO3 │
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 模型切分方式 │ 不切(副本)│ 按 Tensor │ 按层 │ 按参数分片 │
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 显存节省 │ 无 │ 极高(按卡数)│ 高(按层数)│ 极高(按卡数)│
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 通信模式 │ AllReduce │ AllReduce │ P2P(点对点)│ AllGather/ │
│ │ │ │ (每层多次) │ (层间) │ ReduceScatter│
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 通信量/步 │ 中 │ 高 │ 低 │ 高 │
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 计算效率(MFU)│ 高 │ 中 │ 低(有气泡)│ 高 │
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 通信带宽要求 │ 中(IB 可)│ 极高(NVLink)│ 低(IB 可)│ 高(IB 可)│
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 适用场景 │ 显存不够 │ 单层太大 │ 模型太深 │ 显存不够 │
│ ├────────────────┼────────────┼─────────────┼─────────────┼─────────────┤
│ │ 独立可用? │ 可 │ 需要配合 │ 需要配合 │ 可 │
│ │
│ MFU = Model FLOPs Utilization(模型算力利用率) │
│ 业界典型:TF 训练 MFU ≈ 45-60%,理想值 > 50% │
│ │
└─────────────────────────────────────────────────────────────┘实际配置:3D 并行
┌─────────────────────────────────────────────────────────────┐
│ 工业界实际配置:TP + PP + DP 三维并行 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 场景:训练 Llama-3-70B,128 张 H100 │
│ │
│ 配置: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ TP = 8 (节点内张量并行,每层参数切 8 份) │ │
│ │ PP = 8 (32 层分成 8 段,8 个 pipeline stage) │ │
│ │ DP = 16 (数据并行,16 个副本) │ │
│ │ │ │
│ │ 总 GPU 数:8 × 8 × 2 = 128 张 │ │
│ │ (每台机器 8 卡,16 台机器) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 拓扑感知: │
│ → TP = 8:必须同一节点内(NVLink) │
│ → PP:跨节点通信(IB),但通信量小 │
│ → DP:跨节点 AllReduce(IB),通信量大 │
│ │
│ 配置原则: │
│ 1. TP 优先绑 NVLink(带宽最敏感) │
│ 2. PP 其次绑 IB(点对点,IB 够用) │
│ 3. DP 最后填充(IB AllReduce) │
│ │
└─────────────────────────────────────────────────────────────┘升华:并行策略的本质 trade-off
┌─────────────────────────────────────────────────────────────┐
│ 并行策略选择的工程哲学 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 1. 显存和通信的永恒 trade-off │
│ → 切得越碎(TP↑、PP↑):显存占用↓,但通信量↑ │
│ → FSDP 牺牲通信换显存:参数 AllGather → 计算 → ReduceScatter│
│ │
│ 2. 硬件拓扑决定策略上限 │
│ → 没有 NVLink → 不能 TP │
│ → 只有 4 卡/节点 → TP 最多 4 │
│ → IB 带宽不够 → DP 效率低 │
│ │
│ 3. 实际是 TP + PP + DP 的组合优化 │
│ → 没有银弹,每种并行都有自己的适用范围 │
│ → 80B 模型:TP4 + PP4 + DP8(32 卡) │
│ → 400B 模型:TP8 + PP8 + DP16(1024 卡) │
│ │
│ 4. MFU 是最终指标 │
│ → 不是"用了多少卡",而是"卡的实际利用率多少" │
│ → 128 卡 MFU=30% 实际不如 32 卡 MFU=55% │
│ │
│ 一句话总结: │
│ 并行策略 = 在硬件拓扑约束下,找到显存/通信/计算的 Pareto 最优解 │
│ │
└─────────────────────────────────────────────────────────────┘"AI 可查 vs 必须理解"清单
AI 可查:
✅ DeepSpeed / Megatron 的具体配置 YAML
✅ 不同 GPU 集群拓扑下的最优 TP/PP/DP 配置
✅ NCCL 的具体通信原语参数
必须理解:
🔴 数据并行(DDP):每卡完整副本,梯度 AllReduce,简单但显存占用大
🔴 张量并行(TP):单层参数按维度切分,必须 NVLink,通信量极大
🔴 流水线并行(PP):按层切分,存在流水线气泡,需要 1F1B 优化
🔴 FSDP(ZeRO-3):参数分片 + AllGather/ReduceScatter,显存效率高
🔴 为什么 TP 必须同节点(NVLink),PP/DP 可以跨节点(IB)
🔴 3D 并行(TP+PP+DP)是工业界标准配置
🔴 MFU(Model FLOPs Utilization)是衡量并行效率的核心指标学习状态:🟡 开始学习