目录
- 大模型训练的四大核心难点
- [Adam 优化器](#Adam 优化器)
- [解决显存瓶颈:ZeRO 优化器](#解决显存瓶颈:ZeRO 优化器)
- 解决单卡算力不足:数据并行
- 解决模型结构过大(层内拆分):张量并行
- 解决模型结构过大(层间拆分):流水线并行
- 解决流水线气泡:虚拟流水线(VPP)
- [解决 ZeRO-3 跨节点通信瓶颈:ZeRO++](#解决 ZeRO-3 跨节点通信瓶颈:ZeRO++)
- [综合方案:3D 并行](#综合方案:3D 并行)
- 扩展概念
- 选型指南:从模型规模到方案推荐
一、大模型训练的四大核心难点
概述: 训练大模型(7B+)时,会依次遇到显存、计算、通信、工程四个方面的瓶颈。理解这些瓶颈是理解后续所有解决方案的前提。
1.1 显存瓶颈 --- 模型装不进一张卡
一张 A100 (80G) 能装什么?
| 模型 | 参数体积 (fp16) | 训练总需求 | 能否单卡训 |
|---|---|---|---|
| LLaMA-7B | 14 GB | ~56 GB | ✅ 勉强可以 |
| LLaMA-13B | 26 GB | ~104 GB | ❌ |
| LLaMA-70B | 140 GB | ~560 GB | ❌ |
| GPT-3 175B | 350 GB | ~1.4 TB | ❌ |
训练时显存花在哪?
一个 7B 模型训练时的显存分布(fp16 训练 + fp32 Adam):
┌──────────────────────────────────────────────────┐
│ 模型参数 (fp16): 14 GB │
│ 梯度 (fp16): 14 GB │
│ 优化器状态 (fp32 Adam): 56 GB (m+v) │
│ ───────────────────────────────── │
│ 小计: 84 GB │
│ 激活值 (batch=1): 2-4 GB │
│ ───────────────────────────────── │
│ 总计: 86-88 GB │
└──────────────────────────────────────────────────┘
结论: 7B 模型单卡 A100 (80G) 刚好放不下!
↓
这就是为什么需要 ZeRO 分片
计算瓶颈 --- 训练一次等太久
通信瓶颈 --- 卡越多,等通信越久
工程瓶颈 --- 分布式训练有多难
二、Adam 优化器
概述: Adam 是当前大模型训练的事实标准优化器。理解它的工作原理,是理解后面 ZeRO 为什么这么设计的关键------因为 Adam 的显存占用是模型参数的 2 倍。
2.1 优化器的本质
训练的本质:不断调整模型参数 θ,让 loss 最小化。
最简单的 SGD(随机梯度下降):
θ_new = θ_old - lr × gradient
问题:每一步只看当前坡度,没有惯性,震荡大、收敛慢
2.2 Adam 的核心直觉
Adam = 梯度下降 + 动量(惯性) + 自适应步长
| 成分 | 类比 | 作用 |
|---|---|---|
| 动量(Momentum) | 下坡的惯性 | 同方向加速,反方向平滑过渡 |
| 自适应学习率 | 记住抖动幅度 | 抖得厉害的维度步子小,平稳的维度步子大 |
SGD 下山:◉→→→◉↗↘↗↘↗↘◉→→ 震荡大
Adam 下山:◉→→→→→→→→→→→◉→→ 平滑且快
2.3 Adam 的三步计算
# 第1步:计算动量 m(历史梯度的指数加权平均)
m_t = β₁ · m_{t-1} + (1-β₁) · g_t
# ↑ 惯性保留 ↑ 加入新信息
# 第2步:计算方差 v(历史梯度平方的指数加权平均)
v_t = β₂ · v_{t-1} + (1-β₂) · g_t²
# ↑ 历史抖动 ↑ 当前抖动
# 第3步:更新参数
θ_t = θ_{t-1} - lr · m_t / (√v_t + ε)
# 学习率 动量 归一化
2.4 为什么 Adam 这么占显存
每个参数需要额外存两个量(fp32):
模型参数 θ: 1 个值 ← 可以不存?不行
Adam 动量 m: 1 个值 ← 额外
Adam 方差 v: 1 个值 ← 额外
所以优化器状态 = 2 × 参数量 (在 fp32 下)
= 4 × 参数量 (如果模型本身是 fp16)
例子:7B 模型
模型参数 (fp16): 14 GB
优化器状态 (fp32): 56 GB ← 占了总显存的一半!
这就是为什么 ZeRO-1 只分片优化器状态就能省 4x 显存
→ 因为优化器状态本来就和模型参数差不多大
2.5 Adam vs AdamW
# 现在几乎不用 Adam,而是用 AdamW
# AdamW: weight_decay 直接作用在参数更新上,不通过梯度
# (解耦权重衰减 --- Decoupled Weight Decay)
optimizer = torch.optim.AdamW(
model.parameters(),
lr=5e-5,
betas=(0.9, 0.999), # (动量衰减, 方差衰减)
eps=1e-8,
weight_decay=0.01
)
三、解决显存瓶颈:ZeRO 优化器
概述: ZeRO(Zero Redundancy Optimizer)是微软 2020 年提出的显存优化算法,也是 DeepSpeed 的核心。核心思想:不需要每张卡都存全部数据,每张卡只存自己负责的部分。
3.1 核心直觉
传统 DDP:每张卡存完整副本
┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐
│卡0 │ │卡1 │ │卡2 │ │卡3 │
│全量 │ │全量 │ │全量 │ │全量 │ ← 每卡 84GB
└─────┘ └─────┘ └─────┘ └─────┘
总显存: 4 × 84G = 336G(但每卡还是存84G)
ZeRO:分片存储
┌─────┐ ┌─────┐ ┌─────┐ ┌─────┐
│卡0 │ │卡1 │ │卡2 │ │卡3 │
│1/4 │ │1/4 │ │1/4 │ │1/4 │ ← 每卡 21GB
└─────┘ └─────┘ └─────┘ └─────┘
总显存: 4 × 21G = 84G(不冗余了!)
3.2 ZeRO 三级分片
| 阶段 | 分片内容 | 显存节省 | 通信影响 | 一句话 |
|---|---|---|---|---|
| ZeRO-1 | 优化器状态 | 4x | 无额外通信 | 最低成本,推荐默认开启 |
| ZeRO-2 | + 梯度 | 8x | 通信略增 | 推荐大部分场景 |
| ZeRO-3 | + 模型参数 | 线性于卡数 | 通信显著增加 | 超大模型必选 |
3.3 ZeRO-Offload(2021)
解决什么问题: 卡数不够,即使 ZeRO-3 分片后显存还是差一点。
核心直觉: 优化器状态每步只访问一次,可以放到 CPU 内存里。
数据流向:
GPU 显存 CPU 内存
┌─────────┐ ┌──────────┐
│ 参数 │ │ │
│ 梯度 │ │ 优化器状态│ ← 常驻 CPU
│ 激活值 │ │ (m + v) │
└────┬────┘ └────┬─────┘
│ ↑
└── optimizer.step() ─┘
参数拉到 CPU 更新
更新完写回 GPU
| 指标 | ZeRO-3 纯GPU | ZeRO-Offload |
|---|---|---|
| 每卡显存 | 31 GB | 22 GB |
| 训练速度 | 1x | 慢 10-20% |
| 适用场景 | 显存充足 | 显存差一点(10-20%) |
3.4 ZeRO-Infinity(2021)
解决什么问题: 卡极少但模型极大------用硬盘换显存,理论可训练无限大模型。
核心直觉: 训练中每步只访问几层参数,不用的参数可以换出到 NVMe SSD。
三级存储层级:
层级 带宽 容量
GPU HBM 2 TB/s 80 GB ← 每步在用
CPU DRAM 100 GB/s 1-2 TB ← 优化的参数暂存
NVMe SSD 6 GB/s 10-30 TB ← 长期不用的参数
ZeRO-Infinity 动态调度:
每步前:从 CPU/NVMe 加载当前层参数到 GPU
每步后:把不用的参数换出到 CPU/NVMe
类似操作系统的虚拟内存 + 页面置换
| 指标 | ZeRO-3 | ZeRO-Infinity |
|---|---|---|
| 支持最大模型 | 卡数 × 80G | 理论无限 |
| 32卡训 100B | ❌ | ✅ |
| 速度 | 1x | 0.3-0.6x |
| 推荐场景 | 显存够用 | "穷训"超大模型 |
3.5 如何选择 ZeRO Stage
# DeepSpeed 配置中的 ZeRO 选择
# 场景1:显存充裕,只是想加速
"zero_optimization": { "stage": 1 }
# 场景2:通用训练(推荐)
"zero_optimization": { "stage": 2 }
# 场景3:模型超大,必须分片
"zero_optimization": {
"stage": 3,
"contiguous_gradients": true,
"overlap_comm": true
}
# 场景4:显存差一点
"zero_optimization": {
"stage": 3,
"offload_optimizer": { "device": "cpu" } # ZeRO-Offload
}
# 场景5:模型极大但卡极少
"zero_optimization": {
"stage": 3,
"offload_optimizer": { "device": "cpu" },
"offload_param": { "device": "nvme", "nvme_path": "/mnt/nvme" } # Infinity
}
四、解决单卡算力不足:数据并行
概述: 模型装得下了,但一张卡算得太慢------复制多份,每张卡处理不同的数据,吞吐线性增长。
4.1 数据并行的演进
DataParallel(已淘汰)
单进程控制多卡,主卡瓶颈
流程:主卡收数据 → 分发给其他卡 → 等所有卡算完 →
收集梯度 → 主卡更新 → 参数分发给其他卡
问题:主卡的通信量是其他卡的 N-1 倍
N 越大,主卡瓶颈越严重
DistributedDataParallel(DDP,当前标准)
多进程对等多卡,无主卡瓶颈
流程:每张卡独立读数据 → 独立 forward →
独立 backward(自动 all-reduce 同步梯度)→
每张卡独立更新参数
优点:没有主卡,所有卡对等
通信只在 backward 中自动完成
近线性扩展
4.2 DDP + ZeRO 的组合
# DDP 只解决"多卡并行",不解决"显存不够"
# DDP + ZeRO 才是完整方案
# 配置
"zero_optimization": { "stage": 2 }, # ZeRO 解决显存
# DDP 默认开启(多进程本身就是数据并行)
五、解决模型结构过大(层内拆分):张量并行
概述: 当模型的某一层参数太大(如 hidden_size=4096 的 Linear),一张卡算不动或装不下时,把这一层切成多份,每张卡算一份。
5.1 核心直觉
不切(单卡计算一整层):
Linear(4096, 4096) → 单卡算,显存 32MB 参数量
切成2份(两张卡各算一半):
卡0: Linear(4096, 2048) → 16MB
卡1: Linear(4096, 2048) → 16MB
但最后需要 all-reduce 合并结果
5.2 张量并行的代价
| 方面 | 代价 | 说明 |
|---|---|---|
| 通信 | 极高 | 每层都需要 all-reduce |
| 限制 | 仅限节点内 | 依赖 NVLink 高速互联 |
| 典型值 | TP=2 或 TP=4 | 一般不超过 8 |
TP 通信密集的原因:
每层 forward 结束时,需要把各卡的部分结果合并
假设 TP=8:每层一次 all-reduce
96 层 Transformer = 96 次 all-reduce
这只能在 NVLink(600GB/s)上跑
跨以太网(12.5GB/s)完全不可行
六、解决模型结构过大(层间拆分):流水线并行
概述: 当模型层数太多(如 96 层 Transformer),一张卡装不下所有激活值时,按层切段,每张卡负责一段。
6.1 核心直觉
96 层 Transformer
不切:1 张卡存全部 96 层 → 显存爆炸
切成 8 段(PP=8):
卡0: Layer 1-12
卡1: Layer 13-24
卡2: Layer 25-36
...
卡7: Layer 85-96
每张卡只存 12 层的参数和激活值 → 显存省 8x
6.2 1F1B 调度
流水线并行的执行顺序(PP=4, micro_batch=4):
时间 ──────────────────────►
卡0: [F0][F1][F2][F3][B3][B2][B1][B0]
卡1: [F0][F1][F2][F3][B3][B2][B1][B0]
卡2: [F0][F1][F2][F3][B3][B2][B1][B0]
卡3: [F0][F1][F2][F3][B3][B2][B1][B0]
← 这里有空闲
F = Forward(前向)
B = Backward(反向)
问题:最后几张卡在开头空等,前几张卡在结尾空等
这部分空闲比例 = 流水线气泡
6.3 流水线气泡(核心问题)
气泡比例 ≈ (PP-1) / (PP × micro_batch)
示例:
PP=4, micro_batch=4: 气泡 ≈ 19%
PP=8, micro_batch=4: 气泡 ≈ 22%
PP=8, micro_batch=1: 气泡 ≈ 88% ← 严重!
结论:
PP 不能太大(< 8-12 为宜)
micro_batch 不能太小(否则气泡占比太高)
七、解决流水线气泡:虚拟流水线(VPP)
概述: 物理 GPU 数量不变,但通过时间切片让每张卡"看起来像多张卡",交错执行不同微批次的任务,减少空闲等待。
7.1 核心直觉
传统 PP:每张物理卡是 1 个 stage
气泡 ≈ 22%
VPP (VP=2):每张物理卡虚拟成 2 个 stage
卡0 同时处理 微批次A 和 微批次C
卡1 同时处理 微批次B 和 微批次D
交错执行 → 气泡从 22% 降到 12%
VPP (VP=4):每张物理卡虚拟成 4 个 stage
气泡进一步降到 ~6%
7.2 效果对比
| 配置 | 气泡比例 | 吞吐提升 |
|---|---|---|
| 传统 PP (8 stage) | ~22% | 1x(基准) |
| VPP VP=2 | ~12% | +12% |
| VPP VP=4 | ~6% | +18% |
| VPP VP=8 | ~3% | +20% |
// DeepSpeed 配置
{
"pipeline": {
"stages": 8,
"vp": 2, // 虚拟流水线倍数
"activation_partition": true
}
}
八、解决 ZeRO-3 跨节点通信瓶颈:ZeRO++
概述: ZeRO-3 显存省了,但每步需要大量 all-gather/all-reduce 通信。在单机内(NVLink)还能接受,跨节点(以太网)就成了新瓶颈。ZeRO++ 专攻这个问题。
8.1 问题来源
ZeRO-3 的通信开销:
每次 forward: all-gather 参数(从各卡收集参数到当前卡)
每次 backward: all-reduce 梯度(同步梯度到各卡)
8 张卡 × 每步 2 次通信 × 参数量
单机内(NVLink 600GB/s): 通信时间 ~5ms
跨节点(以太网 12.5GB/s): 通信时间 ~250ms ← 比计算还长!
8.2 三个优化
优化一:qgZ --- 梯度量化
# 梯度从 fp16 (2字节) → int8 (1字节)
# 通信量直接减半
# 精度影响:微乎其微(梯度本身有噪声,量化不影响收敛)
优化二:hpZ --- 分层分片
节点内:用 ZeRO-3(全分片,用 NVLink 高速通信)
跨节点:用 ZeRO-1(只分片优化器状态,减少 4x 跨节点通信)
┌──────────┐ ┌──────────┐
│ 节点A │ │ 节点B │
│ ZeRO-3 │ │ ZeRO-3 │
│ ┌─┬─┬─┐ │ │ ┌─┬─┬─┐ │
│ │0│1│2│ │ │ │0│1│2│ │
│ └─┴─┴─┘ │ │ └─┴─┴─┘ │
│ ↑NVLink →│ │ ↑NVLink →│
└────┬─────┘ └────┬─────┘
│ │
└── ZeRO-1 ──────┘
(跨节点只同步优化器)
优化三:cpZ --- 通信与计算重叠
不重叠:
[计算] → [通信] → [计算] → [通信] → ... ← 串行,等通信
重叠:
[计算][通信] ← 计算的同时后台通信
[计算][通信]
通信延迟被计算隐藏,等效"通信不花时间"
8.3 总效果
跨节点通信量对比(8节点 × 8卡 = 64卡):
ZeRO-3 ZeRO++ 减少比例
─────────────────────────────────────────────
跨节点通信 24 GB/step 3 GB/step 87%
训练速度 1x 1.8-2.5x 提升显著
ZeRO++ 适合:
多节点大规模训练(8节点+)
单机多卡场景不需要(NVLink 足够快)
8.4 ZeRO 家族完整对比
| 方案 | 解决什么问题 | 手段 | 显存节省 | 速度影响 | 推荐场景 |
|---|---|---|---|---|---|
| ZeRO-1 | 优化器状态占太多 | 优化器分片 | 4x | 几乎无 | 所有场景默认 |
| ZeRO-2 | + 梯度占太多 | 梯度分片 | 8x | 略慢 | 推荐大部分场景 |
| ZeRO-3 | 参数都放不下 | 参数分片 | 线性 | 增加 | 超大模型 |
| ZeRO-Offload | 卡太少显存差一点 | 卸载到 CPU | 额外 30% | 慢 10-20% | 显存差一口气 |
| ZeRO-Infinity | 卡极少模型极大 | 卸载到 NVMe | 理论无限 | 慢 50-70% | 穷训超大模型 |
| ZeRO++ | 跨节点通信太慢 | 量化+分层+重叠 | --- | 快 2x | 多节点大规模 |
九、综合方案:3D 并行
概述: 单一并行策略都有局限,当模型极大(175B+)时,需要将张量并行 + 流水线并行 + 数据并行三个维度组合使用,取长补短。
9.1 为什么需要三个维度
| 策略 | 局限 | 表现 |
|---|---|---|
| 纯数据并行 | 每卡存完整模型 | 175B 需要 1.4TB → 放不下 |
| 纯张量并行 | 通信太密集 | TP=8 时每层 all-reduce → 跨机跑不了 |
| 纯流水线并行 | 气泡太多 | PP=16 时气泡 ~30% → 算力浪费 |
| 纯 ZeRO-3 | 跨节点通信慢 | 以太网下通信时间 > 计算时间 |
9.2 3D 并行的组合方式
以 GPT-3 175B 在 96 张 A100 上训练为例:
可用卡数: 96
3D 并行配置:
TP=2: 张量并行,每层切 2 份,节点内 NVLink 通信
PP=8: 流水线并行,模型切 8 段,每段 12 层
DP=6: 数据并行,6 个模型副本并行处理不同数据
─────────────────
总卡数 = 2 × 8 × 6 = 96 ✓
一个"DP 组" = 2(TP) × 8(PP) = 16 张卡
这 16 张卡构成一个完整的模型副本
DP=6 表示有 6 个这样的副本
2*8*6=96
9.3 显存分配明细
3D 并行下单卡显存分配(175B 模型):
模型参数: 350GB ÷ (TP=2) ÷ (PP=8) = 22 GB ← 只存自己负责的部分
优化器状态: ZeRO-1 进一步分片到 6 个 DP 组
激活值: 用重计算+较小的 micro_batch 控制
每卡显存 ≈ 22GB + 梯度 + 优化器 + 激活值 ≈ 35-40GB ← A100 80G 装得下
9.4 关键配置参数
{
// 3D 并行三维
"tensor_model_parallel_size": 2,
"pipeline_model_parallel_size": 8,
// ZeRO(每个 DP 组内使用)
"zero_optimization": {
"stage": 1,
"reduce_bucket_size": 5e8
},
// 批量大小计算
"train_micro_batch_size_per_gpu": 4, // 每卡微批次
"gradient_accumulation_steps": 16, // 梯度累积
"train_batch_size": 4 × 2 × 8 × 6 × 16 = 6144 // 总批量
}
十、扩展概念
概述: 在实际大规模训练中,还会用到以下技术作为上述方案的补充。
10.1 混合专家模型(MoE)与专家并行(EP)
MoE 模型结构: 不是所有 token 都经过所有参数,而是路由到不同的"专家"子网络。
MoE 的直觉:
一个 100B 模型不一定是 100B 参数全部激活
可以做成:每个 token 只经过 10B 参数
模型规模变大,但计算量不变 → 用更少的算力做更大的模型
专家并行(Expert Parallel):
MoE 的专家天然可以放在不同的卡上
路由决定哪个 token 去哪个卡
这是 MoE 模型专用的并行策略
| 代表模型 | 参数规模 | 每 token 激活 | 特点 |
|---|---|---|---|
| Switch-Transformer | 1.6T | 约 10B | 每个 token 只走一个专家 |
| Mixtral 8x7B | 47B | 13B | 每个 token 走 2 个专家 |
10.2 序列并行(Sequence Parallelism)
解决什么问题: 长序列训练时,激活值显存随序列长度线性增长,瓶颈在激活值而非参数。
核心思路: 把序列维度也切开,每张卡只存一部分序列的激活值。
传统张量并行切的是 hidden_dim
序列并行额外切 seq_len 维度
适用场景:
长文档理解(16K+ tokens)
多模态大模型(高分辨率图像 = 长视觉序列)
10.3 激活值重计算(Activation Checkpointing / Gradient Checkpointing)
解决什么问题: 不保存中间激活值,反向传播时重新计算------以计算换显存。
forward 时只保存少量关键激活值,删除中间结果
backward 时从这些关键点重新计算中间激活值
显存节省:50-70%
计算额外开销:20-30%
何时使用:
显存不够时用
batch size 上不去时用
10.4 梯度累积(Gradient Accumulation)
解决什么问题: 每卡 batch size 受显存限制,但希望总 batch size 更大(训练更稳)。
# 不累积:每步更新一次参数
for data in loader:
loss = model(data)
loss.backward()
optimizer.step() # batch_size = 4
# 累积 8 步:每 8 步更新一次参数
for i, data in enumerate(loader):
loss = model(data)
loss.backward()
if (i+1) % 8 == 0:
optimizer.step() # 等效 batch_size = 32
optimizer.zero_grad()
10.5 FlashAttention
解决什么问题: 标准 attention 的计算复杂度是 O(n²),且需要大量显存保存中间矩阵。
核心思路: 不让 attention 矩阵完整写出到 HBM,而是在 SRAM 中分片计算。
标准 Attention:
Q×K^T → full matrix → save to HBM → read back → softmax → save → ...
FlashAttention:
分块计算,中间结果不写回 HBM,在 SRAM 内完成
→ 显存省 5-10x,速度快 2-3x
已经集成到 PyTorch 2.0 的 scaled_dot_product_attention
10.6 分布式框架对比
| 框架 | 公司 | 核心能力 | 适合场景 |
|---|---|---|---|
| DeepSpeed | 微软 | ZeRO 发明者,全功能 | 通用训练,ZeRO 首选 |
| FSDP | PyTorch 官方 | 内置 ZeRO-3 | PyTorch 用户,不想加依赖 |
| Megatron-LM | NVIDIA | 张量并行 + 流水线并行 | 超大模型 3D 并行 |
| Megatron-DeepSpeed | 微软+NVIDIA | Megatron + DeepSpeed 合并 | 3D 并行 + ZeRO 的组合 |
| ColossalAI | 潞晨科技 | 多种并行策略集成 | 一站式解决方案 |
| AscendSpeed | 华为 | DeepSpeed 昇腾移植版 | 昇腾 NPU 用户 |
十一、选型指南:从模型规模到方案推荐
概述: 不同规模的模型,适用的方案不同。不要用 175B 的方案去训 7B 的模型------过度设计比不做还糟。
11.1 规模速查表
模型参数 → 推荐方案 ─────────────────────────────────────────────────
< 7B 单卡训(如果能装下)或 纯 DDP
└─ ZeRO-1 可选(省点显存,加速有限)
7B ~ 13B 单机多卡,ZeRO-2
不需要 TP,不需要 PP
└─ 一台 4×A100 或 8×A100 服务器足够
13B ~ 30B 单机多卡,ZeRO-3
可能需要 PP = 4(如果单卡显存不够)
└─ 8×A100 可以训 30B
30B ~ 70B 单机多卡或双机,ZeRO-3 + PP
TP 可选(TP=2 用 NVLink)
└─ 2-4 台 8×A100 服务器
70B ~ 175B 多机多卡,ZeRO-3 + PP + TP
跨节点通信优化(ZeRO++ 可选)
└─ 8-16 台 8×A100 服务器
> 175B 多机多卡,3D 并行(TP+PP+DP/ZeRO)
ZeRO++ 必选(跨节点通信优化)
VPP 可选(减少 PP 气泡)
分布式框架推荐 Megatron-DeepSpeed
└─ 几十到上百台服务器
11.2 决策树
你的模型能装进单卡吗?
│
├─ ✅ 能(< 7B)
│ └─ 需要加速吗?
│ ├─ 不需要 → 单卡训练即可
│ └─ 需要 → DDP + 可选 ZeRO-1
│
└─ ❌ 不能
├─ 卡数足够分片吗?
│ ├─ 够 → ZeRO-2 或 ZeRO-3
│ ├─ 差一点 → ZeRO-Offload
│ └─ 不够 → ZeRO-Infinity
│
└─ 单层参数太大(hidden_size > 单卡容量)?
├─ 是 → + 张量并行 TP
└─ 否 → 跳过
│
└─ 层数太多(激活值超显存)?
├─ 是 → + 流水线并行 PP
│ └─ 气泡高?→ + VPP
└─ 否 → 跳过
│
└─ 需要多节点?
├─ 是 → + ZeRO++(跨节点通信优化)
└─ 否 → 跳过
11.3 实现层面:最小可用配置清单
# 场景:训练 13B 模型,4×A100 (80G)
# 推荐配置:ZeRO-2 + DDP,无需 TP/PP
deepspeed_config = {
"train_batch_size": 128,
"gradient_accumulation_steps": 8,
"fp16": {"enabled": True},
"zero_optimization": {
"stage": 2,
"contiguous_gradients": True,
"overlap_comm": True
},
"optimizer": {
"type": "AdamW",
"params": {"lr": 3e-5, "betas": [0.9, 0.999]}
}
}
# 启动命令
# deepspeed --num_gpus=4 train.py --deepspeed_config ds_config.json
# 场景:训练 70B 模型,8×A100 (80G)
# 推荐配置:ZeRO-3 + PP=4,可选 TP=2
deepspeed_config = {
"train_batch_size": 64,
"gradient_accumulation_steps": 16,
"fp16": {"enabled": True},
"zero_optimization": {
"stage": 3,
"contiguous_gradients": True,
"overlap_comm": True
},
"pipeline": {
"stages": 4,
"vp": 2 # 虚拟流水线,减少气泡
},
"tensor_model_parallel_size": 2,
}
# 启动命令
# deepspeed --num_gpus=8 train.py --deepspeed_config ds_config.json
11.4 昇腾 NPU 环境的对应方案
由于昇腾生态与 CUDA 生态不完全对应,上述方案在昇腾上的映射:
CUDA 方案 昇腾对应方案
────────────────────────────────────────────────
DeepSpeed AscendSpeed(华为移植版)
FSDP torch_npu 内置支持
Megatron-LM mindspore + 手动实现
ZeRO 对应 AscendSpeed 的 zero_optimization
ZeRO-Offload 部分支持(CPU 卸载)
ZeRO-Infinity 暂不支持
ZeRO++ 暂不支持
3D 并行 部分支持(TP+PP 可选,生态不如 CUDA 成熟)
FlashAttention torch_npu 内置 FlashAttention 算子