千卡集群炼大模型:从显存墙到吞吐极限的分布式训练实战指南
目录
- 分布式训练的动机与并行模式总览
- 数据并行DP
- 张量并行TP
- 流水线并行PP
- ZeRO显存优化
- 混合并行与集群拓扑
- 故障恢复与性能调优
- 总结
- 外部引用
摘要
LLM 训练从单卡扩展到千卡集群,核心约束是显存容量与互连带宽。本文沿数据并行、张量并行、流水线并行与 ZeRO 四条主线,量化各方案的单步通信量与显存占用,并给出 DeepSpeed 与 PyTorch 的真实配置。全文以 7B 到 70B 参数规模为例,建立并行策略的选型框架与故障调优路径。
1. 分布式训练的动机与并行模式总览
分布式训练要同时解决两个子问题:把模型装进显存,把算力摊到多卡。模型装不下,靠切分(张量并行、流水线并行、ZeRO 分片);单卡算不完,靠复制(数据并行)。本章先量化单卡瓶颈,再给出三条并行轴的分类,最后定义全文通用的性能标尺。
1.1 显存墙:为什么 7B 模型都装不进一张卡
7B 参数在 FP16 精度下权重占 14GB,这只是起点。训练态的显存需求远不止权重,梯度与优化器状态按相同数量级叠加,单卡方案在 7B 级别即告失效。
- 权重:FP16 每参数 2 字节,7B 参数共 14GB
- 梯度:与权重同尺寸,再占 14GB
- Adam 优化器:fp32 主权重、一阶动量、二阶方差各 28GB,共 84GB
- 训练态合计约 112GB,超过 A100-80GB 的物理显存
- 若把 batch 与序列同时放大,激活值会反超参数显存
激活值是第二道显存墙。序列长度 2048、隐藏维度 4096 时,单层激活缓存约数百 MB,32 层叠加后激活峰值可达 20-40GB。长序列训练下激活占用反超参数,这正是激活重计算与序列并行出现的原因。
| 模型规模 | 参数量 | FP16 权重 | Adam 状态 | 训练态合计 |
|---|---|---|---|---|
| 7B | 7.0B | 14GB | 84GB | 约 112GB |
| 13B | 13.0B | 26GB | 156GB | 约 208GB |
| 70B | 70.0B | 140GB | 840GB | 约 1.1TB |
训练算力同样构成瓶颈。7B 模型训练 2T token 约需 8.4 乘以 10 的 22 次方 FLOPs,单卡 A100 按 45% MFU 计算约需数年,只能靠多卡并行摊平。显存与算力两个约束同时决定了并行是唯一出路。
实际训练还会叠加序列长度与 batch 两个放大项。序列长度从 2048 提到 8192,激活显存约翻四倍,7B 模型的激活峰值从约 20GB 涨到约 80GB,直接逼近单卡上限,这是长上下文训练必须引入激活重计算与序列并行的原因。
| 场景 | 训练态 | 激活峰值 | 单卡总需求 |
|---|---|---|---|
| 7B,seq 2048,batch 1 | 约 112GB | 约 20GB | 约 132GB |
| 7B,seq 8192,batch 1 | 约 112GB | 约 80GB | 约 192GB |
| 13B,seq 2048,batch 1 | 约 208GB | 约 30GB | 约 238GB |
- 激活占比超过训练态时,优先启用激活重计算而不是继续增加分片
- 显存预算先给训练态,余量再给激活,是单卡容量规划的通用顺序
- 分片只能压低参数相关显存,激活显存必须靠重计算与序列并行解决
算力侧同样存在放大项:上下文并行把 attention 的计算量按序列维度切分,FlashAttention 把注意力显存从 O(n2) 降到 O(n),两者都能把单卡放不下的长序列训练拉回可执行范围。显存与算力的交互是选择并行维度的第一判断依据:先算清每参数 16 字节的训练态与激活占比,再决定分片、重计算与并行度的组合。
- 序列长度翻倍时激活显存增长快于计算量,属于典型的放大风险点
- batch 的放大效果与序列相反:batch 翻倍,激活翻倍,但训练态不变
- 显存规划的错误通常在第一周暴露为 OOM 或吞吐塌陷,先做预算再写配置
上述数字把 7B 模型的训练态起步门槛固定在 112GB,说明任何单卡方案都必须先回答显存从哪来。分片解决训练态,重计算与序列并行解决激活,两条路径缺一不可。
- 激活峰值超过训练态时,先重计算后分片,顺序反了会浪费带宽
- 预算数字建议写进训练脚本的断言,配置变更时自动校验
- 长序列与长 batch 是两类风险,分别对应激活与训练态两个维度
1.2 三条并行轴:切数据、切层、切矩阵
三条并行轴切分对象不同,通信粒度差异很大,选型必须同时看显存收益与通信代价。
- 数据并行切训练样本,每卡持完整模型副本,每步一次梯度 AllReduce,通信量与模型规模成正比
- 张量并行切单层内的矩阵乘,每层两次跨卡 AllReduce,依赖 NVLink 高速互连
- 流水线并行按层分组切分,只在 stage 边界传输激活,通信频率低但延迟被放大
- ZeRO 共享数据并行的复制计算,但把优化器、梯度、参数按 rank 分片,按需收集
| 维度 | 切分对象 | 单步通信 | 典型规模 |
|---|---|---|---|
| 数据并行 | 训练样本 | 整模型梯度一次 AllReduce | 16-256 卡 |
| 张量并行 | 矩阵乘行与列 | 每层两次 AllReduce | 2-8 卡 |
| 流水线并行 | 层分组 | 每微批次跨 stage 一次 | 4-16 stage |
| ZeRO | 优化器/梯度/参数 | 分片后按需 AllGather | 256-1024 卡 |
四条路径并不互斥,生产配置几乎总是组合使用。数据并行与 ZeRO 共享同一份复制语义,张量与流水线并行则把模型物理切分。组合时各维度建立独立进程组,通信域互不干扰,总卡数等于各维度乘积。
数据并行与模型并行的根本差异在于通信发生的频率:DP 每步通信一次,TP 每层两次,PP 每微批次一次,ZeRO 每层一次 AllGather。频率越高,对互连延迟越敏感,这正是拓扑分配的出发点。这一频率视角贯穿全文,所有并行方案的优劣最终都能折算成通信频率与单次数据量的乘积,后文每个章节都按这个口径给出量化数字。
- 本节所有通信量都按单步计算,实际调度中按前后向合计计入
- 频率与数据量的乘积是评估并行方案的第一眼判断
1.3 衡量标尺:MFU 与通信占比
MFU(Model FLOPs Utilization)是把吞吐归一化的指标,等于实际算力除以峰值算力。A100 BF16 峰值 312 TFLOPS,训练 7B 时主流集群 MFU 在 35%-55%,张量并行与数据并行重叠做得好才能逼近上限。
- 有效吞吐等于峰值算力乘 MFU,再乘(1 减通信占比)
- 数据并行 64 卡时通信占比约 15%-25%,256 卡时超过 40%
- 张量并行 8 卡内通信占比 5%-10%,跨机后劣化明显
- 流水线并行受气泡率主导,与微批次数量强相关
训练总时长等于总 FLOPs 除以有效吞吐。7B 模型在 1024 卡 A100 且 MFU 为 45% 时约需 7.5 天,这与 LLaMA-2 7B 的公开训练日志量级一致。全文后续章节会反复使用这个量化方法比较不同并行策略。
# 来源:自实现 / budget_estimator.pydefestimate_memory_gb(num_params,dtype_bytes=2):weight_gb=num_params*dtype_bytes/1e9grad_gb=num_params*dtype_bytes/1e9adam_gb=num_params*4*3/1e9returnweight_gb+grad_gb+adam_gbdefestimate_training_days(total_flops,devices,peak_tflops,mfu):achieved=devices*peak_tflops*mfureturntotal_flops/(achieved*1e12)/86400if__name__=="__main__":total_flops=6*7.0e9*2.0e12*2# 6N 乘 token 数days=estimate_training_days(total_flops,1024,312.0,0.45)print(f"7B on 2T tokens:{days:.1f}days with 1024 A100")MFU 的测量要排除 warmup 与评估阶段,取稳态 500-1000 步的均值,否则数字偏低 5-10 个百分点。通信占比的测量用 torch profiler 抓取 NCCL 耗时占比,或用网络计数器的发送字节与理论带宽比对,两个指标配合使用才能判断瓶颈在计算侧还是网络侧。
- 稳态测量是前提,短步数测量会高估 MFU
- 通信占比与 MFU 互补,两者都健康才算配置到位
- 规划时按标称值的 70%-80% 折算,避免乐观估计
- 指标必须在固定 batch 与序列下对比,跨配置比较没有意义
- 标称算力到实际算力之间还有 kernel 效率损耗,MFU 上限按 55% 规划更现实
显存与算力约束确定了三条并行轴的定位,接下来先从实现最朴素、约束最直观的数据并行开始拆解。
2. 数据并行 DP
数据并行是分布式训练的默认起点:复制模型、切分数据、同步梯度。实现最简单,但通信量随模型规模线性增长,显存需求不降反增,属于吞吐维度的手段。