
1. 大模型分布式训练的整体设计思路1.1 为什么单卡训练必然走向分布式先把一个最朴素的问题摆出来一张 80GB 显存的卡能不能训一个 70B 参数的模型答案是不能而且差得很远。我们算一笔账纯推理时 70B 参数如果用 FP16 存储光权重就要 140GB训练时还要加上梯度140GB、优化器状态Adam 的动量和方差FP32 下各 280GB合计 560GB再加上激活值。粗算下来一个 70B 模型全量训练需要的显存轻松突破 1TB。单卡 80GB 连零头都不够。这就是分布式训练存在的根本原因。它不是为了更快这么简单而是不分布式就根本跑不起来。显存墙、算力墙、通信墙三堵墙一起压过来逼着我们把模型切开放到多张卡、多台机器上。但切法有很多种切得不好通信开销能把算力优势全部吃掉。所以理解分布式训练核心不是记住几个名词而是理解每种并行策略到底切了什么、省了什么、又付出了什么通信代价。这篇文章我就按这个逻辑把 DDP、ZeRO、张量并行、流水线并行、上下文并行这几条主线捋清楚配上我实际踩过的坑和可复现的配置。1.2 五种并行策略的分工与边界先给一张全局地图后面每一节再展开。分布式训练的策略可以按切什么来分类策略切分对象主要解决通信模式典型场景DDP数据吞吐量梯度 AllReduce模型能放进单卡ZeRO优化器状态/梯度/参数显存分片通信单卡放不下但层数不多张量并行 TP单层内的矩阵单层显存与算力高频 AllReduce单层巨大节点内流水线并行 PP按层切分模型深度点对点传递层数多跨节点上下文并行 CP序列维度长序列激活环形通信超长上下文这五种不是互斥的实际的大模型训练基本是3D 并行甚至4D 并行——TP 放节点内PP 跨节点DP 兜底再叠加 ZeRO 或 CP。理解它们各自的边界才能组合出合理的方案。提示不要一上来就追求最复杂的组合。我见过太多团队在 8 卡上硬套 4D 并行结果通信开销比省下的显存还贵吞吐反而下降。先跑通 DDP再按瓶颈逐层加策略这是最稳的路径。1.3 选型的第一性原则先定位瓶颈选并行策略之前先问自己三个问题模型能不能放进单卡单层能不能放进单卡序列有多长模型能进单卡只是慢 → 直接 DDP别折腾。模型进不去但单层能进 → ZeRO 或 PP。单层都进不去比如超大 FFN 或超宽 attention→ 必须 TP。序列长到激活爆炸 → 上 CP 或序列并行。这个判断顺序很重要因为它决定了你付出的通信代价从低到高。DDP 通信最少TP 通信最频繁。能少切就少切是分布式训练里最省钱的一条铁律。2. DDP 数据并行最基础也最容易踩坑的一环2.1 DDP 的核心原理与梯度同步机制DDPDistributedDataParallel的思路非常直白每张卡放一份完整的模型副本喂不同的数据 batch各自算梯度然后在反向传播时把所有卡的梯度做一次 AllReduce 求平均保证每张卡的模型始终一致。关键在于什么时候同步。早期 PyTorch 的 DataParallel 是单进程多线程梯度同步在正向结束后统一做效率低。DDP 是多进程每张卡一个进程反向传播时梯度是逐层算出来的DDP 就利用这一点做梯度分桶bucketing把梯度按大小分成若干桶某个桶的梯度全部算完就立刻触发 AllReduce和后续层的反向计算重叠起来。这就是 DDP 能把通信藏进计算里的原因。AllReduce 本身用的是 Ring AllReduce 算法通信量是 2(N-1)/N 倍的参数量N 是卡数。卡越多单卡通信量越接近 2 倍参数量但延迟会上升。所以 DDP 在几十卡以内扩展性很好上百卡就要考虑通信优化了。2.2 DDP 实操配置与关键参数一个最小可用的 DDP 训练脚本骨架我按实际项目里常用的写法给你import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data.distributed import DistributedSampler def setup(): dist.init_process_group(backendnccl) local_rank int(os.environ[LOCAL_RANK]) torch.cuda.set_device(local_rank) return local_rank def main(): local_rank setup() model MyModel().to(local_rank) model DDP(model, device_ids[local_rank], bucket_cap_mb25, # 梯度桶大小 gradient_as_bucket_viewTrue, # 省显存 find_unused_parametersFalse) # 关键默认关掉 sampler DistributedSampler(dataset, shuffleTrue) loader DataLoader(dataset, samplersampler, batch_sizeper_gpu_bs) for epoch in range(epochs): sampler.set_epoch(epoch) # 必须否则每轮打乱顺序一样 for batch in loader: ...启动命令用torchruntorchrun --nproc_per_node8 --nnodes1 \ --master_addr127.0.0.1 --master_port29500 \ train.py这里有几个参数值得单独说。bucket_cap_mb默认 25MB桶越大通信次数越少但重叠越差桶越小重叠越好但通信次数多。实测在 A100 上25MB 到 50MB 之间比较稳具体要看模型层大小。gradient_as_bucket_viewTrue能让梯度直接写进桶里省一份显存拷贝几乎没副作用建议常开。2.3 DDP 常见坑与排查清单DDP 看着简单坑却不少我整理成一张速查表现象可能原因解决卡在 init_process_group端口占用/网络不通换端口检查 NCCL 网卡显存比单卡还高每卡都存了完整优化器状态这是 DDP 固有需上 ZeRO训练变慢且 GPU 利用率低find_unused_parametersTrue确认无未用参数后关掉loss 不收敛忘了 sampler.set_epoch每轮设置 epoch梯度不同步手动 backward 绕过了 DDP hook用 model(batch) 触发find_unused_parametersTrue是我见过最坑的一个参数。它会遍历计算图找未使用参数每步都做开销巨大。如果你的模型确实有条件分支导致部分参数不参与那没办法但绝大多数情况是误开关掉能提速 10% 到 30%。注意DDP 下每张卡都保存完整的优化器状态所以显存占用和单卡几乎一样。DDP 只解决更快不解决更大。想训更大的模型必须往下看 ZeRO。3. ZeRO 系列把优化器状态、梯度、参数逐层切开3.1 ZeRO 的三级切分逻辑ZeROZero Redundancy Optimizer的核心洞察是DDP 里每张卡都存了一份完整的优化器状态、梯度和参数这是巨大的冗余。既然每张卡的数据不同但模型相同那这些状态完全可以切开分给不同的卡需要时再通信拿回来。ZeRO 分三个阶段切得越来越狠ZeRO-1只切优化器状态。显存大头Adam 的 FP32 动量方差被均分到各卡省显存最多通信增加最少。ZeRO-2再切梯度。梯度也算完就扔需要时 AllReduce 回来。ZeRO-3连参数也切。每张卡只存自己那份参数前向反向时按需 AllGather。显存节省大致是ZeRO-1 省约 4 倍优化器状态ZeRO-2 再省梯度ZeRO-3 让显存占用和卡数近似成反比。代价是通信量递增ZeRO-3 的通信量约为 DDP 的 1.5 倍。3.2 DeepSpeed 配置实战ZeRO 最常用的实现是 DeepSpeed。一份 ZeRO-2 的配置长这样{ train_batch_size: 64, gradient_accumulation_steps: 4, zero_optimization: { stage: 2, allgather_partitions: true, allgather_bucket_size: 5e8, overlap_comm: true, reduce_scatter: true, reduce_bucket_size: 5e8, contiguous_gradients: true }, fp16: { enabled: true, loss_scale: 0, initial_scale_power: 16 }, optimizer: { type: AdamW, params: { lr: 1e-4, betas: [0.9, 0.95] } } }几个参数解释一下。overlap_commTrue让通信和计算重叠几乎必开。reduce_bucket_size和allgather_bucket_size控制通信桶大小5e8约 500MB是常见值太小通信频繁太大占显存。contiguous_gradientsTrue让梯度连续存储减少碎片。切到 ZeRO-3 时额外加stage: 3, stage3_prefetch_bucket_size: 5e7, stage3_param_persistence_threshold: 1e5, stage3_max_live_parameters: 1e9, stage3_gather_16bit_weights_on_model_save: truestage3_param_persistence_threshold决定多小的参数不切分比如 LayerNorm 这种小参数留着不切更划算。stage3_gather_16bit_weights_on_model_save在保存时把分片参数聚合成完整权重不加这个存出来的 checkpoint 是碎的。3.3 ZeRO 的显存账与通信代价我拿一个 13B 模型在 8 卡上实测过粗略数据如下FP16 混合精度配置单卡显存相对吞吐DDP约 78GB1.0ZeRO-1约 42GB0.95ZeRO-2约 30GB0.90ZeRO-3约 18GB0.72可以看到 ZeRO-1 几乎白赚显存吞吐损失很小性价比最高。ZeRO-3 省显存最狠但吞吐掉得明显因为参数每次前向都要 AllGather。所以我的经验是能用 ZeRO-1/2 就别上 ZeRO-3除非模型实在太大。实操心得ZeRO-3 配合stage3_prefetch_bucket_size调优能挽回不少吞吐。prefetch 是提前把下一层参数拉过来和当前层计算重叠。这个值设成单层参数大小的一半左右比较合适设太大反而占显存。4. 张量并行与流水线并行切模型的两把刀4.1 张量并行的切分原理张量并行TP切的是单层内部的矩阵运算。以 Transformer 的 FFN 为例第一层是Y XA第二层是Z YB。如果按列切 A按行切 B那么第一层每张卡算X A_i得到部分结果不需要通信。第二层每张卡算Y_i B_i得到部分和需要一次 AllReduce 把结果加起来。这就是 Megatron-LM 提出的经典切法一次前向只需要两次 AllReduceattention 一次FFN 一次。切得巧妙的地方在于它把通信压到了最低同时让每张卡只存一部分权重。但 TP 的通信非常频繁——每一层都要通信而且通信量正比于激活大小。所以 TP 几乎只能放在节点内NVLink 互联跨节点做 TP 会被网络带宽拖死。一般 TP 度数不超过单机 GPU 数比如 8。4.2 流水线并行的气泡问题流水线并行PP按层切分比如 32 层模型切成 4 段每张卡放 8 层。数据像流水线一样从第一段流到最后一段。听起来很美但有个致命问题气泡bubble。如果一次只喂一个 micro-batch那么第一段在算的时候后面几段都在等GPU 利用率极低。解决办法是把一个 batch 切成多个 micro-batch让它们像工厂流水线一样错开填充。这就是 GPipe 和 1F1BOne Forward One Backward调度的由来。气泡占比的近似公式是(P-1)/(MP-1)P 是流水线段数M 是 micro-batch 数。比如 P4、M8气泡占比约 27%M 加到 32气泡降到 8.6%。所以 micro-batch 越多气泡越小但 micro-batch 太多会让单次计算变小通信占比上升。这是个需要调的平衡点。4.3 TP 与 PP 的组合配置实际训练里 TP 和 PP 经常一起用。Megatron-LM 的典型配置是TP8节点内PP4跨节点DP剩余。比如 64 卡可以配成 TP8、PP4、DP2。启动参数大致是torchrun --nproc_per_node8 --nnodes8 \ pretrain_gpt.py \ --tensor-model-parallel-size 8 \ --pipeline-model-parallel-size 4 \ --num-layers 32 \ --hidden-size 4096 \ --num-attention-heads 32 \ --micro-batch-size 4 \ --global-batch-size 512这里micro-batch-size和global-batch-size的关系是global micro × DP × gradient_accumulation。PP 的 micro-batch 数就是 gradient accumulation 步数所以调大 accumulation 能减小气泡。注意TP 和 PP 的通信模式完全不同。TP 是高频 AllReduce必须靠 NVLinkPP 是低频点对点可以走普通网络。所以布局上一定是TP 在节点内PP 跨节点反过来会非常慢。5. 上下文并行长序列训练的新战场5.1 为什么长上下文需要专门的并行当序列长度从 4K 涨到 128K激活值显存是平方级增长的attention 矩阵是 L×L。这时候前面几种并行都不太对症TP 切的是权重不是激活PP 切的是层DDP 每卡都存完整激活。序列维度的激活爆炸需要专门切序列。上下文并行CP就是把序列切成几段每张卡负责一段。难点在于 attention每个 token 要和所有token 算注意力但其他 token 在别的卡上。所以 CP 需要一种特殊的通信模式——环形注意力Ring Attention。5.2 环形注意力的通信机制Ring Attention 的思路很优雅把 KV 也按序列切分每张卡持有自己那段的 Q 和 KV。计算时每张卡用自己的 Q 和本地 KV 算一部分注意力同时把 KV 传给下一张卡接收上一张卡的 KV循环 P 次P 是卡数直到所有 KV 都过了一遍。这样每张卡最终算出了完整的注意力结果而通信和计算可以重叠。通信量是每层 O(L×d)和序列长度线性相关比 attention 本身的 O(L²) 计算量小得多所以重叠得好几乎不增加时间。这也是 CP 能扩展到超长序列的原因。5.3 CP 的适用场景与配置要点CP 不是万能的。它主要解决激活显存对权重显存没帮助。所以 CP 通常和 ZeRO、TP 组合使用。典型场景是长上下文继续预训练或长文本微调。配置上Megatron-LM 和 DeepSpeed 都支持 CP。Megatron 里加--context-parallel-sizeDeepSpeed 里用 Ulysses 或 Ring Attention 实现。要注意的是CP 要求序列长度能被 CP 度数整除且 attention 实现要支持环形通信不是所有 attention 变体都能直接套。实操心得CP 和序列并行SP经常被混淆。SP 是把 LayerNorm、Dropout 这些非 attention 部分的激活也按序列切通常和 TP 绑定使用CP 是专门针对 attention 的序列切分。两者可以叠加但配置时要分清。6. 常见问题与排查技巧实录6.1 通信相关的典型故障分布式训练里通信问题占了故障的一大半。我整理了几类高频问题现象排查方向常用命令NCCL 超时网卡、防火墙、拓扑NCCL_DEBUGINFO卡在某张卡该卡掉队或 OOMnvidia-smi看利用率吞吐忽高忽低通信和计算没重叠好调 bucket size多机跑不起来master 地址/端口检查MASTER_ADDRNCCL_DEBUGINFO是排查通信问题的第一把钥匙它会打印用了哪些网卡、走的什么协议、有没有回退到慢速路径。我遇到过好几次是 NCCL 默认选了错误的网卡加上NCCL_SOCKET_IFNAMEeth0指定网卡就好了。6.2 显存与性能的平衡技巧显存不够时优先级我一般这么排先开 gradient checkpointing用计算换显存省 50% 以上激活再上 ZeRO-1/2然后考虑 TP最后才动 PP 和 CP。gradient checkpointing 是最便宜的省显存手段代价是重算一次前向吞吐掉约 20% 到 30%但比切模型简单太多。性能调优上torch.compile和 FlashAttention 是两个几乎必开的加速项。FlashAttention 把 attention 的显存从 O(L²) 降到 O(L)还更快长序列场景收益巨大。torch.compile能融合算子实测能提速 10% 到 20%但首次编译慢且对动态 shape 支持一般要评估。6.3 从单卡到多卡的迁移检查清单最后给一份迁移清单从单卡脚本改成分布式时逐项核对数据加载是否用了 DistributedSampler且每轮 set_epoch。随机种子是否按 rank 区分避免所有卡数据一样。BatchNorm 是否换成了 SyncBatchNorm如果用了 BN。日志是否只在 rank 0 打印避免刷屏。Checkpoint 是否只在 rank 0 保存或按 ZeRO 分片保存。学习率是否随 global batch 调整线性缩放或 sqrt 缩放。这几条里set_epoch和日志刷屏是最容易被忽略的。前者导致每轮数据顺序一样影响收敛后者在几百卡时能把日志系统打爆。我个人在实际操作中的体会是分布式训练最难的不是写出能跑的代码而是定位瓶颈到底在哪。是显存不够、通信太慢还是计算本身就没喂饱先用 profiler 看清楚时间花在哪再决定加哪种并行比盲目堆策略有效得多。很多时候把 bucket size 调对、把 overlap 打开比换一套并行方案带来的提升还大。