分布式训练这个话题,每次聊起来都让我想起第一次把单卡脚本改成多卡时的那种手忙脚乱。明明在单卡上跑得好好的模型,一上多卡要么显存爆了,要么速度不升反降,要么loss曲线诡异得让人怀疑人生。这篇就来把分布式训练里最常被提到的几条技术路线——DDP、ZeRO、张量并行、流水线并行、上下文并行——从它们各自解决什么问题、底层怎么运作、实际怎么选型,到踩过的坑,系统地捋一遍。不管你是刚接触多卡训练的新手,还是已经在调参一线摸爬滚打的老手,应该都能从中找到对自己有用的东西。
1. 为什么单卡跑不动了:分布式训练的动机与基本盘
1.1 显存墙和算力墙:两个绕不过去的瓶颈
大模型训练最先撞上的墙,几乎都是显存。一个参数量为 N 的模型,训练时显存占用大致可以拆成几块:模型参数本身、梯度、优化器状态(以Adam为例,每个参数需要一阶矩和二阶矩两份状态),再加上前向传播过程中保存的激活值。用混合精度训练时,参数和梯度各占2字节,但优化器状态通常还是FP32,每个参数要占8字节。粗算一下,一个10B参数的模型,光参数+梯度+优化器状态就是 10B × (2+2+8) = 120GB,这还没算激活值。单张80GB的卡根本放不下。
算力墙则是另一个维度的问题。即便显存够用,单卡的计算吞吐也有限。训练一个大模型动辄需要几千甚至上万GPU天的算力,如果只用一张卡,训练周期会长到不可接受。所以分布式训练要同时解决两件事:把显存压力分摊到多张卡上,以及把计算任务并行化来缩短训练时间。
这两件事听起来简单,但做起来涉及一个核心矛盾:卡与卡之间需要通信。通信是有成本的,如果并行策略设计得不好,通信开销会吃掉并行带来的全部收益,甚至让训练比单卡还慢。所以理解分布式训练,本质上是在理解"怎么切分计算和存储,同时把通信控制在可接受范围内"。
1.2 数据并行、模型并行、混合并行:三种切分思路
从切分维度来看,分布式训练的策略可以归为几大类。数据并行是把训练数据切成多份,每张卡持有完整的模型副本,各自算各自的梯度,然后通过通信把梯度做全局平均。这种方式实现简单,是大多数人的入门选择。模型并行则是把模型本身切开,不同的层或同一层内的不同部分放到不同的卡上。模型并行又细分为张量并行(层内切分)和流水线并行(层间切分)。
实际训练大模型时,几乎不会只用一种策略,而是混合并行:比如在节点内用张量并行(因为节点内带宽高),节点间用流水线并行或数据并行,再叠加ZeRO来进一步降低显存。这种组合的思路是让通信量大的并行方式跑在高带宽链路上,通信量小的跑在低带宽链路上。
下面这张表大致概括了几种并行策略的核心特征,方便先建立一个全局印象:
| 策略 | 切分对象 | 主要解决的问题 | 通信特点 | 典型适用场景 |
|---|---|---|---|---|
| 数据并行(DDP) | 数据 | 加速训练、分摊batch | 梯度全规约,通信量大 | 模型能放进单卡 |
| ZeRO | 优化器状态/梯度/参数 | 显存瓶颈 | 按stage不同,通信量递增 | 模型放不进单卡但层数不多 |
| 张量并行 | 层内权重矩阵 | 单层过大 | 每层都要通信,通信频繁 | 节点内高带宽场景 |
| 流水线并行 | 层间 | 模型太深 | 阶段间传激活值,有气泡 | 跨节点、层数多的模型 |
| 上下文并行 | 序列维度 | 长序列激活值爆炸 | 注意力计算需通信 | 超长上下文训练 |
理解了这张表,后面的内容就有了骨架。接下来逐个拆解。
2. DDP:最成熟也最容易踩坑的起点
2.1 DDP到底做了什么:梯度全规约的本质
DDP(DistributedDataParallel)的核心逻辑其实很朴素:每张卡拿到一份完整的模型副本,喂给它不同的数据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(backend="nccl") 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, sampler=sampler, batch_size=per_gpu_bs) for epoch in range(epochs): sampler.set_epoch(epoch) # 这行千万别漏 for batch in loader: ...提示:DDP的通信后端在GPU上要用nccl,gloo只适合CPU或调试。如果初始化时卡住不动,八成是MASTER_ADDR和MASTER_PORT没设对,或者防火墙挡了端口。
2.3 DDP的显存账:为什么它救不了大模型
DDP最大的局限在于:每张卡都要存一份完整的模型、梯度和优化器状态。也就是说,DDP只分摊了数据,没有分摊模型状态。前面算过,10B模型光状态就要120GB,DDP对此无能为力。所以DDP适合的是"模型能放进单卡,但想加速训练"的场景。一旦模型放不进单卡,就必须上ZeRO或模型并行。
另外DDP的通信量也值得注意。梯度全规约的通信量等于梯度总大小,对于大模型来说这个量级不小。虽然Ring AllReduce让它与卡数无关,但绝对量摆在那里,在低带宽的跨节点链路上会成为瓶颈。这也是为什么后来出现了ZeRO——它把优化器状态和梯度也切分开,顺带减少了每张卡需要通信的数据量。
3. ZeRO:把优化器状态、梯度、参数逐级切开
3.1 ZeRO的三个stage:切什么,省多少
ZeRO(Zero Redundancy Optimizer)的核心思想是:DDP里每张卡都存了完整的模型状态,这其实是冗余的。既然数据并行下每张卡算的是不同数据,那能不能把模型状态也切开,每张卡只存一部分,需要的时候再通信拿回来?ZeRO就是干这个的,它分三个递进的stage:
- Stage 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) | 通信量(相对DDP) |
|---|---|---|---|
| 1 | 优化器状态 | 约1/4 | 约1x |
| 2 | +梯度 | 约1/8 | 约1x(用reduce-scatter替代all-reduce) |
| 3 | +参数 | 约1/N | 约1.5x |
3.2 ZeRO-Offload与ZeRO-Infinity:把状态挪到CPU和NVMe
Stage 3 之后还能再省吗?能。ZeRO-Offload把优化器状态和梯度计算卸载到CPU内存,利用CPU做参数更新,GPU只负责前向反向。这样GPU显存进一步释放,代价是CPU和GPU之间的PCIe通信。ZeRO-Infinity更进一步,把状态卸载到NVMe固态硬盘,理论上能训练超出单机内存的模型,但速度受限于存储带宽。
这两个技术适合显存极度紧张、又不想大规模上模型并行的场景。实际用的时候要注意:卸载带来的通信开销可能让训练变慢,需要权衡"能训"和"训得快"。
3.3 ZeRO实践中的坑:通信、配置与stage选择
用DeepSpeed或PyTorch FSDP(Fully 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 Parallelism,TP)。它的思路是把一个大的矩阵乘法拆成多个小矩阵乘法,分到不同卡上算,再把结果拼起来。
以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就需要跨节点,通信瓶颈立刻显现。
这也是为什么实际的大模型训练配置里,经常看到"节点内TP=8,节点间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 Parallelism,PP)是把模型的层按顺序切成若干段,每段放到一张卡(或一组卡)上。数据从第一段流入,逐段处理后从最后一段流出。听起来很直观,但问题在于:如果一次只处理一个batch,那么同一时刻只有一张卡在工作,其他卡都在等,利用率极低。
解决办法是微批次(micro-batch):把一个batch再切成若干小份,让它们像流水线一样依次进入。当第一份数据进入第二段时,第二份数据进入第一段,这样多张卡可以同时工作。但流水线总有"填充"和"排空"的阶段——开始时要等流水线填满,结束时要等它排空,这段时间里部分卡是空闲的,这就是气泡(bubble)。
气泡的大小和流水线深度、微批次数量有关。微批次越多,气泡占比越小。粗略估算,气泡占比约为 (PP度 - 1) / (微批次数 + PP度 - 1)。所以增加微批次数量能有效降低气泡,但微批次太多会让每个微批次的计算量变小,通信占比上升,需要权衡。
5.2 GPipe与1F1B:两种调度策略的取舍
流水线并行有两种经典调度:
- GPipe:先把所有微批次的前向做完,再做所有微批次的反向。这种调度的好处是逻辑简单,但显存占用高——因为要保存所有微批次的激活值直到反向开始。
- 1F1B(One Forward One Backward):前向和反向交替进行,每做完一个微批次的前向就尽快做它的反向,及时释放激活值。显存占用低得多,是现在的主流选择。Megatron和DeepSpeed默认都用1F1B。
1F1B还有变体,比如交错式1F1B(interleaved 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 Parallelism,CP)就是专门针对序列维度的切分。它把输入序列切成若干段,每张卡处理一段。但注意力计算需要每个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组合。比如一个典型的超长上下文训练配置可能是:节点内TP=8处理单层,节点间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卡)配置示例:节点内TP=8,节点间PP=4,DP=2,总共 8×4×2=64 卡。这个配置里TP吃掉节点内带宽,PP和DP走节点间网络。
7.2 通信瓶颈的定位与优化
混合并行下,性能问题往往出在通信上。定位方法:
- 用profiler(如PyTorch Profiler、Nsight Systems)看时间都花在哪,是计算还是通信。
- 检查通信是否和计算重叠。如果通信是串行的,优化空间很大。
- 看网络带宽利用率。如果带宽跑满但速度还是慢,说明通信量本身太大,需要调整并行策略。
优化的方向包括:调整bucket大小、开启通信重叠、把通信量大的并行方式放到高带宽链路、用梯度压缩(如FP16通信、梯度量化)减少通信量。
7.3 常见故障排查表
| 现象 | 可能原因 | 排查方向 |
|---|---|---|
| 训练卡住不动 | 初始化失败、端口冲突 | 检查MASTER_ADDR/PORT、网络连通性 |
| loss不收敛 | 种子不一致、BN问题 | 检查随机种子、归一化层 |
| 速度不升反降 | 通信瓶颈、气泡过大 | profiler定位、调整并行度 |
| OOM | 并行配置不当 | 降stage、开重计算、调bucket |
| 结果不一致 | 通信逻辑错误 | 对比单卡结果、检查all-reduce |
8. 一些踩过的坑和收尾的碎碎念
分布式训练这块,文档里不会写的坑太多了。说几个我印象最深的。
一个是NCCL的环境变量。有时候训练莫名其妙变慢,最后发现是NCCL没走对网卡。NCCL_SOCKET_IFNAME指定网卡、NCCL_IB_DISABLE控制是否用InfiniBand,这些环境变量在多网卡机器上特别关键。默认配置不一定最优,得根据实际硬件调。
另一个是checkpoint的兼容性。不同并行策略下保存的checkpoint格式不一样,换配置后加载经常出问题。建议在训练早期就验证checkpoint的保存和加载流程,别等到训练几天后才发现存下来的东西加载不了。
还有就是小规模验证。上大规模之前,一定先用小模型、小数据在小规模上把并行配置跑通,验证正确性(比如和单卡结果对比)和性能。直接上大规模调试,成本太高。
分布式训练没有银弹,每种并行策略都是在显存、通信、实现复杂度之间做权衡。理解了每种策略解决什么问题、代价是什么,才能根据实际场景做出合理选择。希望这篇能帮你少走点弯路。