
分布式训练这个话题每次聊起来都让我想起第一次把单卡脚本改成多卡时的那种手忙脚乱。明明在单卡上跑得好好的模型一上多卡要么显存爆了要么速度不升反降要么loss曲线诡异得让人怀疑人生。这篇就来把分布式训练里最常被提到的几条技术路线——DDP、ZeRO、张量并行、流水线并行、上下文并行——从它们各自解决什么问题、底层怎么运作、实际怎么选型到踩过的坑系统地捋一遍。不管你是刚接触多卡训练的新手还是已经在调参一线摸爬滚打的老手应该都能从中找到对自己有用的东西。1. 为什么单卡跑不动了分布式训练的动机与基本盘1.1 显存墙和算力墙两个绕不过去的瓶颈大模型训练最先撞上的墙几乎都是显存。一个参数量为 N 的模型训练时显存占用大致可以拆成几块模型参数本身、梯度、优化器状态以Adam为例每个参数需要一阶矩和二阶矩两份状态再加上前向传播过程中保存的激活值。用混合精度训练时参数和梯度各占2字节但优化器状态通常还是FP32每个参数要占8字节。粗算一下一个10B参数的模型光参数梯度优化器状态就是 10B × (228) 120GB这还没算激活值。单张80GB的卡根本放不下。算力墙则是另一个维度的问题。即便显存够用单卡的计算吞吐也有限。训练一个大模型动辄需要几千甚至上万GPU天的算力如果只用一张卡训练周期会长到不可接受。所以分布式训练要同时解决两件事把显存压力分摊到多张卡上以及把计算任务并行化来缩短训练时间。这两件事听起来简单但做起来涉及一个核心矛盾卡与卡之间需要通信。通信是有成本的如果并行策略设计得不好通信开销会吃掉并行带来的全部收益甚至让训练比单卡还慢。所以理解分布式训练本质上是在理解怎么切分计算和存储同时把通信控制在可接受范围内。1.2 数据并行、模型并行、混合并行三种切分思路从切分维度来看分布式训练的策略可以归为几大类。数据并行是把训练数据切成多份每张卡持有完整的模型副本各自算各自的梯度然后通过通信把梯度做全局平均。这种方式实现简单是大多数人的入门选择。模型并行则是把模型本身切开不同的层或同一层内的不同部分放到不同的卡上。模型并行又细分为张量并行层内切分和流水线并行层间切分。实际训练大模型时几乎不会只用一种策略而是混合并行比如在节点内用张量并行因为节点内带宽高节点间用流水线并行或数据并行再叠加ZeRO来进一步降低显存。这种组合的思路是让通信量大的并行方式跑在高带宽链路上通信量小的跑在低带宽链路上。下面这张表大致概括了几种并行策略的核心特征方便先建立一个全局印象策略切分对象主要解决的问题通信特点典型适用场景数据并行DDP数据加速训练、分摊batch梯度全规约通信量大模型能放进单卡ZeRO优化器状态/梯度/参数显存瓶颈按stage不同通信量递增模型放不进单卡但层数不多张量并行层内权重矩阵单层过大每层都要通信通信频繁节点内高带宽场景流水线并行层间模型太深阶段间传激活值有气泡跨节点、层数多的模型上下文并行序列维度长序列激活值爆炸注意力计算需通信超长上下文训练理解了这张表后面的内容就有了骨架。接下来逐个拆解。2. DDP最成熟也最容易踩坑的起点2.1 DDP到底做了什么梯度全规约的本质DDPDistributedDataParallel的核心逻辑其实很朴素每张卡拿到一份完整的模型副本喂给它不同的数据batch各自做前向和反向得到各自的梯度然后通过AllReduce操作把所有卡上的梯度求平均再用平均后的梯度更新参数。因为每张卡的初始参数相同、更新用的梯度也相同所以更新后参数依然保持一致。这里的关键操作是AllReduce。它的目标是把所有卡上的梯度张量逐元素求和或求平均然后把结果广播回每张卡。实现上通常用Ring AllReduce算法所有卡排成一个环分两个阶段——先做reduce-scatter把梯度分块每块在一部分卡上完成累加再做all-gather把累加好的块广播回所有卡。这个算法的好处是通信量与卡数无关只与梯度总量有关所以扩展性很好。DDP还有一个重要优化叫梯度分桶bucketing。反向传播是从后往前逐层计算梯度的如果每算完一层就通信一次通信次数会非常多而且每次通信的数据量小链路利用率低。DDP的做法是把多个梯度张量打包成一个bucket等bucket填满再触发一次AllReduce同时通信和反向计算可以重叠overlap进一步隐藏通信开销。2.2 从单卡到DDP改造脚本时最容易忽略的几件事把单卡脚本改成DDP表面上看只是加几行初始化代码但实际有几个地方特别容易出问题。第一是随机种子。如果每张卡的随机种子不一样数据增强、dropout这些随机操作会产生不同的结果虽然梯度平均后大体还能收敛但会引入额外的噪声。正确做法是给每张卡设置不同的种子用于数据shuffle的差异化但模型初始化相关的种子要保证一致或者干脆在初始化后从rank 0广播参数。第二是BatchNorm。DDP默认每张卡独立计算BN统计量如果每卡batch size很小BN的统计会很不稳定。这时候要么改用SyncBN跨卡同步统计量但通信开销大要么换成LayerNorm/GroupNorm这类不依赖batch维度的归一化。第三是数据加载器的分布式采样。必须用DistributedSampler否则每张卡会读到重复数据等于变相减小了有效batch size。而且要注意在每轮epoch开始时调用sampler.set_epoch(epoch)否则shuffle的随机性在每轮之间会重复。# DDP初始化的典型写法 import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data.distributed import DistributedSampler dist.init_process_group(backendnccl) local_rank int(os.environ[LOCAL_RANK]) torch.cuda.set_device(local_rank) model MyModel().to(local_rank) model DDP(model, device_ids[local_rank]) sampler DistributedSampler(dataset) loader DataLoader(dataset, samplersampler, batch_sizeper_gpu_bs) for epoch in range(epochs): sampler.set_epoch(epoch) # 这行千万别漏 for batch in loader: ...提示DDP的通信后端在GPU上要用ncclgloo只适合CPU或调试。如果初始化时卡住不动八成是MASTER_ADDR和MASTER_PORT没设对或者防火墙挡了端口。2.3 DDP的显存账为什么它救不了大模型DDP最大的局限在于每张卡都要存一份完整的模型、梯度和优化器状态。也就是说DDP只分摊了数据没有分摊模型状态。前面算过10B模型光状态就要120GBDDP对此无能为力。所以DDP适合的是模型能放进单卡但想加速训练的场景。一旦模型放不进单卡就必须上ZeRO或模型并行。另外DDP的通信量也值得注意。梯度全规约的通信量等于梯度总大小对于大模型来说这个量级不小。虽然Ring AllReduce让它与卡数无关但绝对量摆在那里在低带宽的跨节点链路上会成为瓶颈。这也是为什么后来出现了ZeRO——它把优化器状态和梯度也切分开顺带减少了每张卡需要通信的数据量。3. ZeRO把优化器状态、梯度、参数逐级切开3.1 ZeRO的三个stage切什么省多少ZeROZero Redundancy Optimizer的核心思想是DDP里每张卡都存了完整的模型状态这其实是冗余的。既然数据并行下每张卡算的是不同数据那能不能把模型状态也切开每张卡只存一部分需要的时候再通信拿回来ZeRO就是干这个的它分三个递进的stageStage 1只切分优化器状态。每张卡只存 1/N 的优化器状态N为卡数更新参数时通过通信收集。显存节省约4倍针对优化器状态部分。Stage 2在Stage 1基础上再切分梯度。每张卡只存 1/N 的梯度反向传播时通过reduce-scatter直接得到分片的梯度。显存进一步节省。Stage 3再切分模型参数。每张卡只存 1/N 的参数前向和反向时按需通过all-gather把参数收集回来。显存节省最大但通信量也最大。用一个具体例子感受一下假设模型有 7.5B 参数用Adam混合精度训练。DDP下每卡需要 7.5B × 16字节 ≈ 120GB。ZeRO Stage 1 把优化器状态8字节/参数切到8张卡上每卡省下约 7.5B × 8 × (7/8) ≈ 52GB。Stage 2 再切梯度Stage 3 再切参数逐级把单卡占用压到能放进80GB甚至更小的卡里。Stage切分内容单卡显存相对DDP通信量相对DDP1优化器状态约1/4约1x2梯度约1/8约1x用reduce-scatter替代all-reduce3参数约1/N约1.5x3.2 ZeRO-Offload与ZeRO-Infinity把状态挪到CPU和NVMeStage 3 之后还能再省吗能。ZeRO-Offload把优化器状态和梯度计算卸载到CPU内存利用CPU做参数更新GPU只负责前向反向。这样GPU显存进一步释放代价是CPU和GPU之间的PCIe通信。ZeRO-Infinity更进一步把状态卸载到NVMe固态硬盘理论上能训练超出单机内存的模型但速度受限于存储带宽。这两个技术适合显存极度紧张、又不想大规模上模型并行的场景。实际用的时候要注意卸载带来的通信开销可能让训练变慢需要权衡能训和训得快。3.3 ZeRO实践中的坑通信、配置与stage选择用DeepSpeed或PyTorch FSDPFully Sharded Data Parallel本质是ZeRO-3的工程实现时有几个坑很典型。坑一stage选太高反而慢。很多人一上来就开Stage 3结果发现训练速度掉了一大截。原因是Stage 3每层前向都要all-gather参数反向又要重新收集通信极其频繁。如果模型其实能放进单卡用Stage 1甚至DDP就够了。选stage的原则是能用低stage解决显存问题就别上高stage。坑二通信和计算的重叠没配好。DeepSpeed里有个overlap_comm参数开启后能让通信和计算重叠。但重叠需要额外的显存做缓冲显存紧张时开了反而OOM。这个要结合实际情况调。坑三checkpoint保存和加载。ZeRO-3下每张卡只存了部分参数保存checkpoint时需要特殊处理比如DeepSpeed的stage3_gather_16bit_weights_on_model_save否则存下来的权重是不完整的。加载时也要对应配置不然会报形状不匹配。// DeepSpeed ZeRO-2 的典型配置片段 { zero_optimization: { stage: 2, offload_optimizer: { device: cpu }, overlap_comm: true, contiguous_gradients: true, reduce_bucket_size: 5e8, allgather_bucket_size: 5e8 }, fp16: { enabled: true } }注意reduce_bucket_size和allgather_bucket_size这两个参数直接影响通信效率。设太小通信次数多设太大显存占用高。一般从5e8开始调显存不够就往下调。4. 张量并行把一层拆到多张卡上4.1 张量并行的切分逻辑列切与行切当单层权重矩阵大到一张卡放不下时数据并行和ZeRO都帮不上忙它们不切分层内结构这时候需要张量并行Tensor ParallelismTP。它的思路是把一个大的矩阵乘法拆成多个小矩阵乘法分到不同卡上算再把结果拼起来。以Transformer里的线性层 Y XW 为例W的形状是 [输入维度, 输出维度]。张量并行有两种切法列并行把W按输出维度切成N份每张卡算 Y_i X W_i得到输出的一个列块。这种切法每张卡都需要完整的输入X输出是拼接关系。行并行把W按输入维度切成N份每张卡算 Y_i X_i W_i得到部分和最后需要AllReduce把各部分加起来。Megatron-LM的设计很巧妙在Transformer的注意力块里QKV投影用列并行输出投影用行并行这样中间不需要额外的同步MLP里第一个线性层用列并行第二个用行并行同样让通信只在必要的地方发生。这种列切接行切的配对能把每层的通信次数压到最低。4.2 张量并行为什么必须待在节点内张量并行的通信特点是每一层的前向和反向都要通信。一个几十层的Transformer意味着几十次甚至上百次通信。如果这些通信跑在跨节点的低速链路上开销会大到无法接受。所以张量并行几乎总是限制在单个节点内利用NVLink或NVSwitch这种高带宽互连带宽可达数百GB/s甚至更高。节点内的卡数通常是8张所以张量并行度一般不超过8。超过8就需要跨节点通信瓶颈立刻显现。这也是为什么实际的大模型训练配置里经常看到节点内TP8节点间PP或DP的组合。TP负责把单层拆开PP负责把层拆开DP负责把数据拆开各司其职。4.3 TP的实操细节通信原语与性能调优实现张量并行时核心用到的通信原语是AllReduce行并行的部分和和AllGather列并行的输出拼接有时可以省掉。Megatron-LM里还用了f和g两个算子来标记前向和反向的通信点框架会自动插入对应的通信操作。调优方面几个经验点通信与计算重叠TP的通信可以和相邻层的计算重叠Megatron里通过调整通信算子的位置来实现。配置得当能隐藏相当一部分通信时间。序列并行Sequence Parallelism这是TP的一个补充把LayerNorm和Dropout这些逐元素操作也按序列维度切开进一步降低激活值显存。它和TP配合使用在Megatron里是标配。避免频繁的小通信如果TP度设得很大每张卡算的矩阵很小通信占比就会飙升。一般TP度不超过8超过就要考虑换策略。提示TP对模型代码有侵入性需要改写线性层和注意力的实现。如果不想改代码可以用Megatron-LM或DeepSpeed的TP支持但灵活性会受限。5. 流水线并行按层切分与气泡的博弈5.1 流水线并行的基本模型阶段、微批次与气泡流水线并行Pipeline ParallelismPP是把模型的层按顺序切成若干段每段放到一张卡或一组卡上。数据从第一段流入逐段处理后从最后一段流出。听起来很直观但问题在于如果一次只处理一个batch那么同一时刻只有一张卡在工作其他卡都在等利用率极低。解决办法是微批次micro-batch把一个batch再切成若干小份让它们像流水线一样依次进入。当第一份数据进入第二段时第二份数据进入第一段这样多张卡可以同时工作。但流水线总有填充和排空的阶段——开始时要等流水线填满结束时要等它排空这段时间里部分卡是空闲的这就是气泡bubble。气泡的大小和流水线深度、微批次数量有关。微批次越多气泡占比越小。粗略估算气泡占比约为 (PP度 - 1) / (微批次数 PP度 - 1)。所以增加微批次数量能有效降低气泡但微批次太多会让每个微批次的计算量变小通信占比上升需要权衡。5.2 GPipe与1F1B两种调度策略的取舍流水线并行有两种经典调度GPipe先把所有微批次的前向做完再做所有微批次的反向。这种调度的好处是逻辑简单但显存占用高——因为要保存所有微批次的激活值直到反向开始。1F1BOne Forward One Backward前向和反向交替进行每做完一个微批次的前向就尽快做它的反向及时释放激活值。显存占用低得多是现在的主流选择。Megatron和DeepSpeed默认都用1F1B。1F1B还有变体比如交错式1F1Binterleaved 1F1B把模型切成更多段并让每张卡负责多个不连续的段进一步降低气泡。代价是通信次数增加。调度策略显存占用气泡大小实现复杂度GPipe高较大低1F1B低中等中交错1F1B低小高5.3 PP的工程实践阶段划分与负载均衡PP落地时最头疼的是负载均衡。如果各段的计算量不均最慢的那段会成为瓶颈其他段都在等它。Transformer里各层结构相同按理说均分就好但第一段和最后一段通常还要承担embedding和输出层计算量和显存占用都不同需要特殊处理。实践中常见的做法是把embedding和输出层单独算或者给它们分配更少的层。Megatron里可以指定每张卡放几层手动调平衡。另外PP的通信量相对小只在阶段边界传激活值所以适合跨节点部署和TP形成互补。还有一个细节是激活值重计算activation recomputation。PP下每张卡要保存自己那段的激活值用于反向如果段内层数多激活值显存会很大。开启重计算后前向时不保存中间激活反向时重新算一遍用计算换显存。这个技术在TP和PP里都很常用。6. 上下文并行长序列训练的新战场6.1 长上下文为什么需要专门的并行策略当序列长度从几千涨到几万甚至几十万时激活值的显存占用会线性甚至平方级增长。注意力的计算复杂度是 O(L²)激活值里注意力矩阵就占了大头。这时候光靠TP、PP、ZeRO都不够因为它们在序列维度上没有切分。上下文并行Context ParallelismCP就是专门针对序列维度的切分。它把输入序列切成若干段每张卡处理一段。但注意力计算需要每个token看到序列里的其他token所以切分后必须通过通信交换信息。6.2 环形注意力与通信模式CP的核心难点在注意力。以Ring Attention为代表的方案把序列切成N段分到N张卡上每张卡持有自己那段的Q、K、V。计算时K和V像环一样在各卡之间流转每张卡用本地的Q和流转过来的K、V算一部分注意力逐步累加。这样每张卡最终得到完整的注意力输出而显存只需要存自己那段的激活。通信量方面Ring Attention的通信量和序列长度、卡数相关但通过和计算重叠可以把通信隐藏在注意力计算背后。这也是它相比其他方案的优势。6.3 CP与其他并行策略的组合CP通常不单独使用而是和TP、PP、DP组合。比如一个典型的超长上下文训练配置可能是节点内TP8处理单层节点间PP处理层数CP处理序列DP处理数据。这种多维并行的配置非常复杂需要框架支持如Megatron-LM的CP支持、DeepSpeed的序列并行。实际调的时候要注意CP的切分粒度、通信重叠的配置、以及和TP的交互。CP和TP都涉及注意力如果两者叠加通信模式会更复杂需要仔细验证正确性。7. 混合并行实战怎么组合怎么选7.1 一个可参考的配置思路实际训练大模型时并行策略的组合没有标准答案但有一个大致的决策顺序先看模型能不能放进单卡。能就用DDP或ZeRO-1加速。放不进单卡但层数不多。用ZeRO-2或ZeRO-3配合offload。单层特别大比如超大hidden size。上张量并行限制在节点内。层数特别多。上流水线并行跨节点部署。序列特别长。上上下文并行。以上组合形成TP×PP×CP×DP的多维并行。一个常见的8节点64卡配置示例节点内TP8节点间PP4DP2总共 8×4×264 卡。这个配置里TP吃掉节点内带宽PP和DP走节点间网络。7.2 通信瓶颈的定位与优化混合并行下性能问题往往出在通信上。定位方法用profiler如PyTorch Profiler、Nsight Systems看时间都花在哪是计算还是通信。检查通信是否和计算重叠。如果通信是串行的优化空间很大。看网络带宽利用率。如果带宽跑满但速度还是慢说明通信量本身太大需要调整并行策略。优化的方向包括调整bucket大小、开启通信重叠、把通信量大的并行方式放到高带宽链路、用梯度压缩如FP16通信、梯度量化减少通信量。7.3 常见故障排查表现象可能原因排查方向训练卡住不动初始化失败、端口冲突检查MASTER_ADDR/PORT、网络连通性loss不收敛种子不一致、BN问题检查随机种子、归一化层速度不升反降通信瓶颈、气泡过大profiler定位、调整并行度OOM并行配置不当降stage、开重计算、调bucket结果不一致通信逻辑错误对比单卡结果、检查all-reduce8. 一些踩过的坑和收尾的碎碎念分布式训练这块文档里不会写的坑太多了。说几个我印象最深的。一个是NCCL的环境变量。有时候训练莫名其妙变慢最后发现是NCCL没走对网卡。NCCL_SOCKET_IFNAME指定网卡、NCCL_IB_DISABLE控制是否用InfiniBand这些环境变量在多网卡机器上特别关键。默认配置不一定最优得根据实际硬件调。另一个是checkpoint的兼容性。不同并行策略下保存的checkpoint格式不一样换配置后加载经常出问题。建议在训练早期就验证checkpoint的保存和加载流程别等到训练几天后才发现存下来的东西加载不了。还有就是小规模验证。上大规模之前一定先用小模型、小数据在小规模上把并行配置跑通验证正确性比如和单卡结果对比和性能。直接上大规模调试成本太高。分布式训练没有银弹每种并行策略都是在显存、通信、实现复杂度之间做权衡。理解了每种策略解决什么问题、代价是什么才能根据实际场景做出合理选择。希望这篇能帮你少走点弯路。