千卡集群炼大模型:从显存墙到吞吐极限的分布式训练实战指南

目录

  1. 分布式训练的动机与并行模式总览
  2. 数据并行DP
  3. 张量并行TP
  4. 流水线并行PP
  5. ZeRO显存优化
  6. 混合并行与集群拓扑
  7. 故障恢复与性能调优
  8. 总结
  9. 外部引用

摘要

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 状态训练态合计
7B7.0B14GB84GB约 112GB
13B13.0B26GB156GB约 208GB
70B70.0B140GB840GB约 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 三条并行轴:切数据、切层、切矩阵

LLM Training

Data Parallel

Tensor Parallel

Pipeline Parallel

ZeRO Sharding

Replicate Full Model

Split One Layer

Split Layer Groups

Shard Optimizer and Params

AllReduce Gradient Per Step

AllReduce Activation Per Layer

Transfer Activation Between Stages

Gather Parameters On Demand

三条并行轴切分对象不同,通信粒度差异很大,选型必须同时看显存收益与通信代价。

  • 数据并行切训练样本,每卡持完整模型副本,每步一次梯度 AllReduce,通信量与模型规模成正比
  • 张量并行切单层内的矩阵乘,每层两次跨卡 AllReduce,依赖 NVLink 高速互连
  • 流水线并行按层分组切分,只在 stage 边界传输激活,通信频率低但延迟被放大
  • ZeRO 共享数据并行的复制计算,但把优化器、梯度、参数按 rank 分片,按需收集
维度切分对象单步通信典型规模
数据并行训练样本整模型梯度一次 AllReduce16-256 卡
张量并行矩阵乘行与列每层两次 AllReduce2-8 卡
流水线并行层分组每微批次跨 stage 一次4-16 stage
ZeRO优化器/梯度/参数分片后按需 AllGather256-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

数据并行是分布式训练的默认起点:复制模型、切分数据、同步梯度。实现最简单,但通信量随模型规模线性增长,显存需求不降反增,属于吞吐维度的手段。

2.1 DDP 的训练循环与梯度同步