分布式训练框架对比
分布式训练框架对比
大模型大到一张 GPU 放不下时怎么办?答案是拆开——把模型或数据拆到多张卡上一起算。怎么拆、用什么工具拆,就是本文要讲的事。
一、先搞清楚问题:为什么一张 GPU 不够?
以 LLaMA-70B 为例,算一下一张 H100 80GB 能不能跑:
模型参数: 70B × 2 bytes (FP16) = 140 GB ← 已超 80GB,别说训练了
梯度: 70B × 2 bytes = 140 GB
优化器状态: 70B × 4 × 3 (Adam) = 840 GB ← fp32 param + momentum + variance
激活值(batch=1, seq=4096): ≈ 50 GB
─────────────────────────────────────────────────
训练总共需要: ≈ 1170 GB
单卡显存: 80 GB
缺口: 1090 GB ← 需要大约 15 张 H100
结论:大模型训练天生就是分布式问题,必须多卡协作。
二、三种基本的「拆法」
2.1 数据并行(DP)—— 一人算一份数据,最后对答案
3 张 GPU,每张有一个完整模型副本:
GPU 0: [完整模型] 喂 batch 0 → 算出梯度 A
GPU 1: [完整模型] 喂 batch 1 → 算出梯度 B
GPU 2: [完整模型] 喂 batch 2 → 算出梯度 C
然后三人对答案:AllReduce → 求平均梯度 → 各自更新模型
结果:三张卡的模型始终保持一致
类比:三个学生做不同卷子,对答案后统一下次做题策略。
适用:模型本身能放进单卡(参数 + 优化器 < 显存),但想加速训练。
局限:每张卡都要存完整模型。LLaMA-70B 单卡参数就要 140GB,DP 根本跑不了——得用下面两种。
2.2 张量并行(TP)—— 把每一层切成几块,分给不同 GPU
原始:一层 Attention 在 GPU 0 上完整计算
┌───────────────────────────────┐
│ QKV投影 → Attention → 输出投影 │ ← GPU 0 单干
└───────────────────────────────┘
TP=2:同一层切成两半
┌──────────────────┐ ┌──────────────────┐
│ QKV投影(前一半) │ │ QKV投影(后一半) │
│ Attention(前一半) │ │ Attention(后一半) │
│ 输出投影(前一半) │ │ 输出投影(后一半) │
└──────────────────┘ └──────────────────┘
GPU 0 GPU 1
↑──── 每步都要通信 ────→ 交换中间结果
类比:一道 100 行的矩阵乘法,A 算前 50 行,B 算后 50 行,算完后拼起来。
特点:切得越细,通信越频繁。所以 TP 只能在 NVLink 全互联的节点内用(8 卡 DGX/HGX),不能跨节点——PCIe 带宽扛不住。
适用:单层参数太大,一张卡算不完。
2.3 流水线并行(PP)—— 前半截模型在一组 GPU,后半截在另一组
模型有 24 层 Transformer:
GPU 0-3: Layer 0-7 (前半截) ────→ GPU 4-7: Layer 8-15 ────→ GPU 8-11: Layer 16-23
↑ 每个阶段之间只传递激活值(很小),通信开销极低 ↑
类比:工厂流水线——A 车间做毛坯,传给 B 车间精加工,B 传给 C 车间组装。每个车间只负责一段。
特点:通信量最小(只传激活值),跨节点友好。但流水线有空泡——前一个阶段没算完,后一个阶段只能等着。
适用:层数深的大模型,跨节点带宽有限时优先用 PP 而非 TP。
三、三种策略怎么组合?
真实训练中很少只用一种,而是混搭:
以 GPT-175B + 1024 张 A100 为例的 3D 并行:
第一维 TP=8(节点内):
└── 把每层切成 8 块 → 单节点 8 卡 NVSwitch 全互联
第二维 PP=8(跨节点):
└── 整个模型分成 8 段 → 8 个节点串成流水线
第三维 DP=16(全局):
└── 上面那个 8×8=64 GPU 的配置复制 16 份 → 16 份数据同时训练
8 × 8 × 16 = 1024 GPU
TP 解决单层放不下的问题,PP 解决层数太多的问题,DP 解决数据太多的问题。三者各司其职。
四、框架的本质:帮你实现上面这些拆法
不同框架就是不同等级的「拆模型工具箱」:
| 框架 | 一句话 | 适合什么时候用 |
|---|---|---|
| PyTorch DDP | 只做数据并行,不改模型代码 | 模型能放进单卡(<10B),单纯想加速 |
| PyTorch FSDP | DDP 升级版:自动分片参数,省显存 | 模型 10-70B,不想改代码,用 PyTorch 原生方案 |
| DeepSpeed | FSDP 的竞品,Microsoft 出品 | 同上,想要更多配置选项和 offload 能力 |
| Megatron-LM | NVIDIA 出品,精细控制 TP+PP+DP | 100B+ 模型,追求极致吞吐,愿意重写模型 |
| ColossalAI | 社区方案,功能多但不稳定 | 实验阶段 |
| TorchTitan | Meta 的 FSDP 最佳实践参考 | 学习 FSDP 用法,不推荐直接用于生产 |
运维视角的关键差异:
框架侵入性(越小越好改/迁移):
DDP(无) < FSDP/DeepSpeed(低) < TorchTitan(中) < Megatron(高)
需要运维掌握的程度:
DDP(只需 torchrun) < FSDP(加几个参数) < DeepSpeed(加配置文件)
< Megatron(管拓扑/NVSwitch/IB,参数巨多)
五、DeepSpeed ZeRO:自动省显存的 DP
DeepSpeed 最大的创新是 ZeRO——在数据并行的基础上,自动分片存储优化器状态、梯度和参数,极大节省显存:
DDP: 每张卡存完整的 参数 + 梯度 + 优化器(Adam m/v)
ZeRO-1:优化器状态分片(每卡存 1/N)→ 省 ~4×
ZeRO-2:梯度也分片 → 省 ~8×
ZeRO-3:参数也分片(用时才从其他卡"借")→ 省 ~N×,但通信多 50%
类比:DDP 是每人家里存全套百科全书,ZeRO-3 是小区共享图书馆——你需要哪页就去借,用完还回去。省地方,但借还有时间成本。
六、选框架的实操指南
你的情况 推荐
─────────────────────────────────────────────────────
单卡能放下模型,想加速 → DDP(零改动,直接 torchrun)
多卡但模型 ≤ 70B,不想改代码 → FSDP 或 DeepSpeed ZeRO-3
多卡,模型 > 100B,追求极致性能 → Megatron-LM 3D 并行
你是算法工程师,不想管分布式细节 → DeepSpeed ZeRO-3(一个 JSON 配完)
你是集群运维,要帮算法调性能 → Megatron-LM(控制力最强,但也最复杂)
MoE 模型(如 Mixtral) → DeepSpeed + Expert Parallel
用 AMD GPU → PyTorch FSDP(DeepSpeed 对 ROCm 支持弱)
七、运维实战:启动命令速查
DDP(最简单)
torchrun --nproc_per_node=8 --nnodes=2 \
--node_rank=$RANK --master_addr=$MASTER --master_port=29500 \
train.py
FSDP(一行参数开启)
torchrun --nproc_per_node=8 --nnodes=4 \
--node_rank=$RANK --master_addr=$MASTER --master_port=29500 \
train.py --fsdp "full_shard auto_wrap" --bf16
DeepSpeed ZeRO-3(一个 JSON + deepspeed 命令)
// ds_config.json 关键配置
{ "zero_optimization": { "stage": 3 } }
deepspeed --num_gpus=8 --num_nodes=4 \
--master_addr=$MASTER --master_port=29500 \
train.py --deepspeed_config ds_config.json
Megatron-LM 3D 并行(参数最多,控制最细)
torchrun --nnodes=64 --nproc_per_node=8 \
pretrain_gpt.py \
--tensor-model-parallel-size 4 \ # TP=4,节点内
--pipeline-model-parallel-size 4 \ # PP=4,节点间
# DP 自动 = 64×8/(4×4) = 32
--bf16 --use-flash-attn
八、常见问题
| 问题 | 最可能原因 | 排查方向 |
|---|---|---|
| GPU 利用率低(< 80%) | 通信瓶颈,TP 跨节点了 | 确认 TP 只在同节点内 |
| ZeRO-3 特别慢 | CPU offload 拖后腿 | 减少 offload 或加 GPU |
| 多节点训练 OOM | batch size 太大或 TP 切分不对 | 先单节点调通再加节点 |
| DeepSpeed 初始化失败 | NCCL_IB_HCA 配错 | 见 ../network/NCCL 通信原理与调优 |
关联知识
- PyTorch 分布式训练实战 — 手把手写 FSDP 训练脚本
- ../network/NCCL 通信原理与调优 — 通信是怎么跑的
- ../performance/GPU 集群性能调优指南 — 调 MFU
- ../hardware/NVLink 与 NVSwitch 拓扑详解 — TP 为什么只能在节点内
- ../GPU 集群运维知识总览 — 返回总览
参考资源
学习时间
| 阶段 | 时间 | 备注 |
|---|---|---|
| 初版框架 | 2026-06-29 | 骨架 |
| 重写 | 2026-06-30 | 降低门槛,以问题和场景驱动 |
状态标记
📖 已掌握 — 三种并行策略的本质区别、框架选型决策、ZeRO 省显存原理 📝 待补充 — 各框架实际 benchmark 数据 (MFU)