专栏大模型笔记·并行1 / 5
21 min学习

并行策略总览:六个维度各切什么、怎么通信

「并行」章的开篇。一次训练里能下刀的位置只有六个:batch、序列、权重矩阵的行列、层、expert,以及注意力内部的序列。每种切法一节,回答同一组问题:切什么、为什么能切、前向和反向各发生什么通信、通信量多大、每卡的计算和显存少了多少。最后看 Megatron、DeepSpeed、PyTorch DTensor / FSDP / torchtitan、JAX GSPMD 怎么用各自的方式表达同一套切法,以及它们怎么在一个 device mesh 上叠起来。

目录11 节

一步训练要做的事是固定的:一个 batch 进来,前向算出 loss,反向算出每个参数的梯度,优化器更新参数。单卡放不下、或者单卡太慢,就要把这一步拆到很多卡上。能拆的方式并不多:训练涉及的每个张量就那么几个维度,每一种并行就是沿其中一个维度下刀。刀落在哪,就决定了每张卡少算什么、少存什么,以及原本在一张卡里免费完成的数据搬运变成了什么通信。

这一篇把六种切法一个一个讲。每一节回答同一组问题:切的是什么为什么可以这样切(哪条数学恒等式保证结果不变),前向和反向各发生了什么通信通信量多大每张卡的计算和显存各少了多少。讲完六种再看主流框架怎么表达它们:Megatron-LM 手写并行层,PyTorch 的 DTensor / FSDP2 / torchtitan 用「张量放置」自动推出通信,JAX 的 GSPMD 让编译器来做这件事。三种表达背后是同一套切法。

先把要切的东西摆出来

记模型有 LL 层、隐藏维 ddaa 个注意力头;一个 micro-batch 有 bb 条序列、每条 SS 个 token;一个 batch 分成 mm 个 micro-batch。一层 dense transformer 的权重约 12d212d^2 个参数(注意力 4d24d^2,MLP 8d28d^2),整个模型 Ψ12Ld2\Psi \approx 12Ld^2。一层的前向 FLOPs 约 24bSd2+4bS2d24\,bSd^2 + 4\,bS^2d,前一项是矩阵乘,后一项是注意力分数;反向是前向的两倍。显存有两块:参数相关的静态部分,混合精度加 Adam 是每参数 16 字节(ZeRO 那篇有这笔账);激活是动态部分,每层约 sbd(34+5as/d)sbd\,(34 + 5as/d) 字节,这是 Megatron 的激活重计算论文的公式,后面序列并行一节会用到。

层 1层 2PP:在层之间切展开一层:激活,权重(序列)条序列DP:切 batchSP / CP:切序列TP:切列或切行ZeRO / FSDP:只切、梯度、优化器状态的存储,计算前 all-gather 凑齐MoE 层的个 expertexpert 1expert 2expertEP:在 expert 之间切FFN 若是 MoE
一次训练里可以下刀的六个位置。DP 切 batch,SP 和 CP 切序列,TP 切权重矩阵的行或列,PP 切层,EP 切 expert;ZeRO / FSDP 不改变计算的切法,只把权重、梯度、优化器状态的存储按 DP 组切开。后面每一节讲其中一刀。

图里六条虚线就是六种并行。它们分成三类。切 batch 和切序列是切数据:每张卡拿到的样本或 token 少了,模型不变。切权重的行列、切层、切 expert 是切模型:每张卡只有一部分参数。注意力内部切序列是切数据里比较特殊的一种,因为注意力是唯一让 token 之间相互看的算子,切了序列就得把别人的 KV 搬过来。ZeRO / FSDP 不在这六条线里:它不改变任何计算的切法,只把参数、梯度、优化器状态的存储按数据并行组切开,用的时候再凑齐。

通信都由五个集合操作组成。记一个张量在一张卡上的完整大小是 MM 字节,nn 张卡参与:

操作做什么每卡收发量
all-reduce所有卡的 MM 逐元素相加,每张卡都拿到结果2M\approx 2M
reduce-scatter相加,但每张卡只拿结果的 1/n1/nM\approx M
all-gather每张卡出 M/nM/n,拼成完整的 MMM\approx M
all-to-all每张卡把自己的 MM 切成 nn 份,第 jj 份发给卡 jjM\approx M
点对点一张卡发给另一张卡MM

ring 算法下 all-reduce 的精确量是 2(n1)M/n2(n-1)M/n,all-gather、reduce-scatter、all-to-all 是 (n1)M/n(n-1)M/nnn 大时就是表里的近似。要记住的只有一件事:all-reduce = reduce-scatter + all-gather。这个等式在 ZeRO、序列并行、FSDP 里反复出现。

数据并行:切 batch

切什么。 batch。NN 张卡,每张卡拿 b/Nb/N 条序列,跑一份完整的模型。

为什么可以。 loss 是样本的平均,L=1bxbatch(x;W)\mathcal{L} = \frac{1}{b}\sum_{x\in\text{batch}} \ell(x; W),梯度是线性算子,所以

WL=1Ni=1NNbxshardiW(x;W)gi.\nabla_W \mathcal{L} = \frac{1}{N}\sum_{i=1}^{N} \underbrace{\frac{N}{b}\sum_{x \in \text{shard}_i}\nabla_W\ell(x; W)}_{g_i} .

每张卡各算自己分片上的梯度 gig_i,平均一下就是整个 batch 的梯度。这要求 loss 对样本可加,transformer 里没有 BatchNorm 之类跨样本的算子,所以严格成立。

DDP:复制权重,切 batchGPU 0条序列完整一份前向、反向,本地GPU 1条序列完整一份前向、反向,本地all-reduce每步通信一次,量个数,与 batch 无关FSDP / ZeRO-3:权重也切开,用前凑齐GPU 0条序列常驻,虚框是用时凑齐的算出完整,只留自己的GPU 1条序列常驻,虚框是用时凑齐的算出完整,只留自己的all-gatherreduce-scatter每层前向、反向各 all-gather 一次,反向再 reduce-scatter,量
左:DDP。每卡一份完整的 ,各算自己 batch 分片的梯度,一次 all-reduce 之后两卡拿到同一个 ,做同一个更新。右:FSDP / ZeRO-3。 只存 ,前向和反向用到某一层时 all-gather 凑齐、用完即丢;梯度 reduce-scatter 后每卡只留自己那 。两种形态的计算完全一样,差别只在 的存储。

前向。 没有通信。每张卡是一个完整的模型,各算各的。

反向。 每张卡算出完整的 gig_i,做一次 all-reduce 得到 gˉ\bar g。之后每张卡用同一个 gˉ\bar g 更新同一份 WWNN 份权重始终一致。实现上不会等反向全部结束再通信:梯度是按层从后往前算出来的,算完一桶就发一桶(PyTorch DDP 的 bucket),通信藏在剩余层的反向后面。

通信量。 每步一次 all-reduce,Ψ\Psi 个梯度,BF16 下 4Ψ\approx 4\Psi 字节。这个量和 batch 大小无关:batch 越大、每步算得越久,通信占比就越小。这是数据并行能跨节点、能到几千卡的原因。

计算与显存。 计算完美地分掉 1/N1/N。显存是数据并行的短板:每张卡一份完整的 WWgg、优化器状态,16 字节乘 Ψ\Psi,7B 模型就是 112 GB,NN 张卡上有 NN 份完全相同的东西。激活因为 batch 变小而减少到 1/N1/N

分片形态。 ZeRO 和 FSDP 补的就是这个短板,图右半是它的终点形态。优化器状态、梯度、参数依次按 NN 切开,每张卡常驻 1/N1/N;前向或反向走到某一层时 all-gather 出那一层的完整权重,用完丢掉;反向的梯度不做 all-reduce 而做 reduce-scatter,每张卡只留自己那一段。计算完全没变,只是 WWgg 的存储换了地方,代价是每层多一次 all-gather,通信从 2Ψ2\Psi 涨到 3Ψ3\Psi。三级各省多少、通信怎么变,在 ZeRO 那篇里展开。

张量并行:切权重矩阵的行和列

切什么。 一层里的一个矩阵。Megatron-LM 的做法是把 MLP 的两个矩阵、注意力的四个矩阵各切成 TT 份,每张卡拿一份,TT 张卡合起来算一层。切的维度是 dd4d4d,也就是权重矩阵的行或列,所以它也叫模型并行里的「层内并行」。

为什么可以。 分块矩阵乘。MLP 是 Y=GeLU(XA)BY = \mathrm{GeLU}(XA)\,BARd×4dA \in \mathbb{R}^{d\times 4d}BR4d×dB \in \mathbb{R}^{4d\times d}。有两种切法,选哪种取决于中间那个非线性。

AA 按列切成 [A1 A2][A_1\ A_2],那么 XA=[XA1  XA2]XA = [XA_1\ \ XA_2]:每张卡拿完整的 XX 和一半的列,算出一半的中间结果。GeLU 是逐元素的,GeLU([XA1  XA2])=[GeLU(XA1)  GeLU(XA2)]\mathrm{GeLU}([XA_1\ \ XA_2]) = [\mathrm{GeLU}(XA_1)\ \ \mathrm{GeLU}(XA_2)],所以每张卡对自己那一半做 GeLU,不需要通信。反过来如果 AA 按行切,XA=X1A1+X2A2XA = X_1A_1 + X_2A_2 是一个和,GeLU 不能穿过加号,做非线性之前就得先 all-reduce 一次。这是 Megatron 选列切的全部理由。

BB 接在 GeLU 后面,此时每张卡手里是 Hi=GeLU(XAi)H_i = \mathrm{GeLU}(XA_i),正好是 HH 的一半列。把 BB 按行切成 [B1;B2][B_1; B_2],那么

Y=[H1  H2][B1B2]=H1B1+H2B2.Y = [H_1\ \ H_2]\begin{bmatrix} B_1 \\ B_2\end{bmatrix} = H_1B_1 + H_2B_2 .

每张卡算出一个部分和 Yi=HiBiY_i = H_iB_i,一次 all-reduce 相加。列切接行切,整个 MLP 只在最后通信一次。

两卡各一份GPU 0GeLU逐元素部分和,GPU 1GeLU逐元素部分和,两卡各一份all-reduce:前向恒等,反向 all-reduce:前向 all-reduce,反向恒等  两者互为反向按列切(,各卡算一半的列按行切(部分和要相加,这就是 all-reduce
Megatron 的 MLP 切法。 按列切成 ,每卡算 ,逐元素的非线性不跨列,所以中间不通信; 按行切成 ,每卡算出一个部分和 ,一次 all-reduce 相加。 是一对共轭算子:前向 恒等、 all-reduce,反向反过来。一层 MLP 前向一次、反向一次 all-reduce,注意力子层同理,所以每层每个 micro-batch 共四次。

注意力是同一个套路,切的单位是头。Q,K,VQ, K, V 的投影矩阵按列切,让每张卡拿到 a/Ta/T 个头的完整 Q,K,VQ, K, V,softmax 在头内部做,所以注意力本身不需要通信。输出投影 WOW_O 按行切,和 BB 一样,最后一次 all-reduce。

前向和反向。 图里的 ffgg 是 Megatron 定义的一对共轭算子。ff 在进入并行区时:前向恒等(每张卡本来就有完整的 XX),反向 all-reduce(两张卡对 XX 的梯度要相加)。gg 在离开并行区时:前向 all-reduce,反向恒等(YY 的梯度直接给每张卡一份)。所以一个 MLP 子层前向一次 all-reduce、反向一次;注意力子层同样;一层每个 micro-batch 共四次 all-reduce,每次的对象是一个 b×S×db\times S\times d 的激活。

通信量。 每层每个 micro-batch 约 4×2bSd4 \times 2\,bSd 个数 =16bSd= 16\,bSd 字节(BF16)。和数据并行比一比:数据并行每步通信 4Ψ=48Ld24\Psi = 48Ld^2 字节,张量并行每步 16bSdLm16\,bSd \cdot L \cdot m 字节,比值是 bSm/3dbSm / 3d。一个 DP rank 每步处理的 token 数 bSmbSm 通常是几万到几十万,dd 是几千,所以张量并行的通信量是数据并行的十倍以上,而且分散在每一层、每次都要等。这就是它必须放在节点内 NVLink 上的原因,也是 TT 一般不超过 8 的原因。

计算与显存。 矩阵乘的 FLOPs 分掉 1/T1/T。参数、梯度、优化器状态的大头(六个矩阵)也是 1/T1/T。但有两样东西没有切:一是 LayerNorm、dropout、残差相加,它们在每张卡上重复算,计算量小,但它们的输入激活每张卡一份,这部分占每层激活的 10sbd10\,sbd 字节,TT 再大也不减少;二是 XXYY 本身也是复制的。用前面的公式,只开 TP 的每层激活是 sbd(10+24/T+5as/(dT))sbd\,(10 + 24/T + 5as/(dT)),那个不随 TT 缩小的 10 是下一节的动机。还有一个实现细节:残差路径上的 dropout 在 TT 张卡上重复执行,必须用同一个随机种子,否则各卡的 XX 就不一致了,Megatron 为此维护了一套模型并行的随机数状态。

序列并行:把张量并行切不到的部分按序列切

切什么。 上一节剩下的那 10sbd10\,sbd:LayerNorm、dropout、残差的输入。它们都是逐 token 或逐元素的算子,一个 token 的输出只依赖这个 token 自己,所以可以沿序列切成 TT 份,每张卡只处理 S/TS/T 个 token,和张量并行用同一组卡。这是 Megatron 2022 年论文里的序列并行(sequence parallelism),DeepSpeed-Ulysses 和 ring attention 也常被叫作序列并行,但它们切的是注意力内部,本文放在上下文并行一节。

为什么可以。 两个事实。第一,LayerNorm 对每个 token 独立归一化,dropout 对每个元素独立采样,残差是逐元素加,沿序列切不改变任何一个 token 的结果。第二,all-reduce = reduce-scatter + all-gather。进入张量并行区之前,激活是按序列切的,每张卡有 S/TS/T 个 token,要做的是 all-gather 拼出全部 SS 个 token(图里的 fˉ\bar f);离开张量并行区时,每张卡手里是全部 SS 个 token 的部分和,本来要 all-reduce,现在换成 reduce-scatter:相加,但每张卡只拿自己那 S/TS/T 个 token 的结果(gˉ\bar g)。原来的一次 all-reduce 被拆成了两半,一半放在进口,一半放在出口。

一层 transformer,张卡的张量并行组;上行注意力子层,下行 MLP 子层注意力子层MLP 子层每卡持有的激活每卡持有的激活LayerNorm逐 tokenall-gather反向 reduce-scatter注意力按头切,每卡个头(子层内部)reduce-scatter反向 all-gatherdropout + 残差逐元素LayerNorm逐 tokenall-gather反向 reduce-scatterMLP切列,切行(子层内部)reduce-scatter反向 all-gatherdropout + 残差逐元素不切时每层激活;只用 TP:,其中 10 那份(虚线框的输入)在张卡上是复制的加上 SP 之后全部除以
序列并行(Megatron-SP)。虚线框里的 LayerNorm、dropout、残差按序列切,每卡只有 个 token;实线框里的注意力和 MLP 按头 / 按列切,每卡有全部 token 但只有 的通道。进 TP 区 all-gather(),出 TP 区 reduce-scatter(),两者合起来正好是原来的一次 all-reduce,通信量不变,但原本复制 份的那部分激活也切开了。

前向和反向。 fˉ\bar f 前向 all-gather、反向 reduce-scatter;gˉ\bar g 前向 reduce-scatter、反向 all-gather。一层前向两次 all-gather 加两次 reduce-scatter,反向同样。

通信量。 一次 all-gather 加一次 reduce-scatter,恰好等于一次 all-reduce。所以通信量和纯张量并行完全一样16bSd16\,bSd 字节每层每 micro-batch,只是操作换了形式。序列并行是白拿的。

计算与显存。 LayerNorm、dropout、残差的计算不再重复,每卡只做 1/T1/T。激活从 sbd(10+24/T+5as/(dT))sbd\,(10 + 24/T + 5as/(dT)) 变成 sbd(34/T+5as/(dT))sbd\,(34/T + 5as/(dT)):全部激活都被 TT 除了。在 T=8T = 8 时,这一步把每层激活砍掉约一半,是不用重计算就能省激活的最大一笔。

流水线并行:切层

切什么。 层。LL 层切成 PP 段,每段连续 L/PL/P 层放一张卡(一个 stage)。切的维度是层的序号,不是任何一个张量的内部。

为什么可以。 层之间是函数复合,y=fLf1(x)y = f_L \circ \cdots \circ f_1(x),把它拆到 PP 张卡上执行不需要任何恒等式。问题恰恰在于它太容易了:第 kk 段必须等第 k1k-1 段算完才能开始,PP 张卡里同一时刻只有一张在干活。让它们同时忙起来的唯一办法是引入互相独立的工作:把 batch 切成 mm 个 micro-batch,第 kk 段做 micro-batch jj 的时候,第 k1k-1 段做 micro-batch j+1j+1。这就是流水线。所以流水线并行天然要求 micro-batch,也天然有一段没人可以干活的时间:第一个 micro-batch 流到最后一段之前,和最后一个 micro-batch 的反向流回第一段之前。

GPU 0,stage 0层 1层 2层 3参数,只有这GPU 1,stage 1层 4层 5层 6参数,只有这GPU 2,stage 2层 7层 8层 9参数,只有这GPU 3,stage 3层 10层 11层 12参数,只有这激活梯度点对点GPipe 时间线,,前向 1 格、反向 2 格。斜线:气泡stage 0stage 1stage 2stage 3F1F2F3F4F1F2F3F4F1F2F3F4F1F2F3F4B1B2B3B4B1B2B3B4B1B2B3B4B1B2B3B4时间灌满排空每个 stage 忙,总长,气泡占比前向反向
流水线并行。上: 层切成 段,段间只传一个 的激活(前向)和同样大小的激活梯度(反向),是所有并行里通信最少的。下:GPipe 时间线, 个 micro-batch,前向 1 格、反向 2 格。斜线是气泡:stage 3 开头要等前向流到它,stage 0 结尾要等反向流回来。1F1B、虚拟段、Zero Bubble 怎么处理它,见流水线那篇。

前向和反向。 前向时第 kk 段把输出激活(b×S×db\times S\times d)点对点发给第 k+1k+1 段。反向时第 k+1k+1 段把对这个激活的梯度(同样大小)发回第 kk 段。每个 micro-batch 在每个段边界上一来一回,没有集合通信。反向需要前向时的激活,所以每个 stage 要为所有「前向已做、反向未到」的 micro-batch 存着激活。

通信量。 每个 micro-batch 在每个边界 2bSd×2=4bSd2\,bSd \times 2 = 4\,bSd 字节,P1P-1 个边界。和张量并行比:张量并行是每层 16bSd16\,bSd,流水线是每段(L/PL/P 层)4bSd4\,bSd,相差 4L/P4L/P 倍。这是流水线能跨节点的原因。

计算与显存。 每张卡只算 L/PL/P 层。但时间不是简单地除以 PP:GPipe 那种「所有前向做完再做所有反向」的排法,总时长是 (m+P1)(tf+tb)(m + P - 1)(t_f + t_b),其中只有 m(tf+tb)m(t_f+t_b) 在干活,气泡占比 (P1)/(m+P1)(P-1)/(m+P-1)。参数、梯度、优化器状态干净地分成 1/P1/P。激活取决于调度:GPipe 下每个 stage 要存全部 mm 份,1F1B 下 stage rrPrP-r 份,前面的 stage 存得最多。气泡怎么算、四代调度各拿什么去换它,是流水线那篇的内容。

expert 并行:切 expert

切什么。 MoE 层的 expert。一层有 EE 个 expert,RR 张卡每张放 E/RE/R 个,非 expert 部分(注意力、router、shared expert)每张卡一份。切的维度是 expert 的编号。

为什么可以。 MoE 层对每个 token 的输出是

y=etop-k(x)ge(x)Ee(x),y = \sum_{e \in \text{top-}k(x)} g_e(x)\, E_e(x) ,

每个 expert 是一个独立的 FFN,只作用在被路由到它的那些 token 上。把一个 micro-batch 的全部 token 按目标 expert 重新排列,同一个 expert 的 token 排在一起,每个 expert 就是一次普通的矩阵乘;算完再按原来的顺序排回去。排列和逆排列是精确的,不改变任何结果。当 expert 分布在不同卡上,这个排列就是一次跨卡的 all-to-all。

EP rank 0router 之后的本地 token00010203颜色 = 目标 expertexpert 0FFN,权重只在本卡0011expert 1FFN,权重只在本卡0213回到原位,加权求和00010203EP rank 1router 之后的本地 token10111213颜色 = 目标 expertexpert 2FFN,权重只在本卡0110expert 3FFN,权重只在本卡0312回到原位,加权求和10111213dispatch:all-to-all每 token 发份,共个数combine:all-to-all原路返回,同样个数虚线是跨 rank 的那部分:个 rank 时平均的 token 要出卡。收到 token 最多的 rank 决定这一层的时间。
expert 并行。 个 expert 切到 个 rank 上,每个 token 的颜色是 router 选中的 expert。dispatch 是一次 all-to-all:把 token 送到持有它目标 expert 的 rank;每个 expert 对收到的 token 做一次普通的 FFN;combine 再一次 all-to-all 把结果送回原位,按 router 权重求和。两次 all-to-all 合起来是一个置换,反向把它们的方向调转即可。

前向和反向。 前向两次 all-to-all:dispatch 把每个 token 发给它的 kk 个目标 expert 所在的卡,combine 把 kk 个结果收回来加权求和。反向也是两次:combine 的反向是把 yy 的梯度按同样的路径发出去,dispatch 的反向是把各 expert 对输入的梯度收回来。all-to-all 的转置还是 all-to-all,方向反过来就行。

通信量。 一个 token 是 dd 个数,发 kk 份,dispatch 是 kbSdk\,bSd 个数,combine 同样,前向 4kbSd4k\,bSd 字节、反向同样。k=8k = 8 时每层 64bSd64\,bSd比张量并行还多,而且 RR 通常大于 8,一定要跨节点。这是 MoE 训练里 all-to-all 成为焦点的原因:把 dispatch 压成 FP8、和 shared expert 的计算重叠、按拓扑分两跳发送,都是在对付这个数。不过这个量和 RR 几乎无关:RR 张卡时平均 (R1)/R(R-1)/R 的 token 要出卡,RR 大了就接近全部。

计算与显存。 expert 的参数分掉 1/R1/R,这是 MoE 模型能做到几万亿参数的原因:Kimi K3 有 2.7T 的 routed expert 参数,靠 EP 切到几百张卡上。非 expert 部分复制,用数据并行处理。计算上每张卡算它收到的 token,收到最多的那张卡决定这一层的时间:router 不会把 token 均匀分给 expert,一张卡的 expert 热了,其他卡就要等它。处理这个偏斜有几条路:给 expert 设容量上限、超出的 token 丢掉(早期做法);在 loss 里加负载均衡项让 router 学着均匀;给热 expert 在别的卡上放副本。最后一条路在 Kimi K3 的 MoonEP里推到了极致,每张卡最多 E/RE/R 个冗余 expert 就能做到完全平衡。

还有一个和张量并行的对比值得说:expert 的矩阵也可以用张量并行切,但一个 expert 本来就比 dense 层的 FFN 窄得多,再切 TT 份,矩阵乘就小到跑不满 GPU;按 expert 切保留了每个矩阵乘的完整形状。这是 MoE 模型偏爱 EP 而不是 TP 的算力上的理由。

上下文并行:在注意力内部切序列

切什么。 序列,但这次切进注意力里面。CC 张卡每张拿 S/CS/C 个 token 的 Q,K,VQ, K, V,各自算这些 token 的注意力输出。前面序列并行绕开了注意力,因为注意力是唯一让 token 之间互相看的算子;上下文并行正面处理它。它的动机是 SS 长到一张卡放不下这条序列的激活,5as2b/d5as^2b/d 那一项随 SS 平方增长。

为什么可以。 softmax 注意力的输出对每一个查询 qq

o=jesjvjjesj,sj=qkj/dh.o = \frac{\sum_j e^{s_j} v_j}{\sum_j e^{s_j}} , \qquad s_j = q^\top k_j / \sqrt{d_h} .

分子和分母都是对 jj 的求和,可以分块累加,唯一的麻烦是数值稳定要减去最大值,而最大值要看完全部 jj 才知道。在线 softmax(也是 FlashAttention 的核心)解决了它:维护三个量,当前见过的最大值 mm、分母 ll、未归一化的分子 OO。收到一块新的 Kj,VjK_j, V_j,先算这一块自己的 (mj,lj,Oj)(m_j, l_j, O_j),再合并:

m=max(m,mj),l=lemm+ljemjm,O=Oemm+Ojemjm.m' = \max(m, m_j), \quad l' = l\,e^{m - m'} + l_j\,e^{m_j - m'}, \quad O' = O\,e^{m - m'} + O_j\,e^{m_j - m'} .

合并顺序任意,全部合并完 O/lO/l 就是精确的注意力输出。所以每张卡只需要依次看到所有 KV 块,不需要同时拿到。ring attention 就是让 KV 块在 CC 张卡之间轮转:第 tt 步 rank ii 手里是 rank (it)modC(i - t) \bmod C 的 KV,算一块、并一块、把这块传给下家、从上家收下一块。CC 步走完每个 rank 都见过全部 KV。

KV 块沿环轮转,;每格是这一步 rank 手里的 KV 块第 0 步第 1 步第 2 步第 3 步rank 0持有,累积rank 1持有,累积rank 2持有,累积rank 3持有,累积每步收一块、发一块,个数,与计算重叠步走完,就是完整注意力的输出连续切分zigzag 切分格子数 3 / 7 / 11 / 15格子数 9 / 9 / 9 / 9查询行是查询块,列是键块,颜色是负责的 rank;因果掩码只算下三角rank:连续切分拿第段;zigzag 拿第和第
上下文并行(ring attention)。 个 rank,rank 持有 不动,KV 块沿环每步传一格:第 步 rank 拿到的是 步后每个 见过所有 KV。每一步用在线 softmax 把新块并进 ,结果与一次算完全一样。右:因果掩码下只有下三角的块要算,按连续段切时 rank 3 的活是 rank 0 的四倍;zigzag 把序列切成 段、rank 拿第 和第 段,每个 rank 的格子数一样多。

前向和反向。 前向每步发出一块 K,VK, V、收进一块,同时算手里的一块,通信和计算重叠。反向要重新走一遍环:每张卡要对所有 KV 块算 dK,dVdK, dV 的贡献,这些贡献属于别的 rank,所以反向除了轮转 K,VK, V 还要轮转累积中的 dK,dVdK, dV

通信量。 每步一块 KV,(S/C)2d(S/C)\cdot 2d 个数,走 C1C - 1 步,前向每张卡收发约 2bSd2=4bSd2\,bSd \cdot 2 = 4\,bSd 字节,CC 无关;反向约三倍。但每步的计算量是 (S/C)2d(S/C)^2 d 级别,随 CC 平方下降,通信固定、计算变小,CC 大到一定程度通信就藏不住了。GQA 会让这个数好看很多:轮转的是 KV 头,GQA 的 KV 头数只有查询头的几分之一。另一条路是 DeepSpeed-Ulysses:不轮转 KV,而是在注意力前后各做一次 all-to-all,把「按序列切」换成「按头切」,注意力内部每张卡有全部 token 的 a/Ca/C 个头,通信量随 CC 减小,代价是 CC 不能超过头数。线性注意力是第三条路:它的「KV」是一个固定大小的状态矩阵,跨卡只需传状态而不是 token,通信量与 SS 无关,Kimi K3 的 KCP 是这个做法。

计算与显存。 激活干净地分掉 1/C1/C,包括注意力那一项。参数不切,每张卡一份,所以它总和别的并行叠用。计算有一个因果掩码带来的坑:只算下三角,按连续段切的话最后一个 rank 要算的块最多,第一个 rank 几乎没活。图右边的 zigzag 切法把序列分成 2C2C 段,rank ii 拿第 ii 段和第 2C1i2C-1-i 段,一头一尾配对,每个 rank 的块数相等。Megatron(经 TransformerEngine)和 PyTorch 的上下文并行都用这类头尾配对的切法。

框架怎么表达这六种切法

上面六节里,每种并行都是「某个张量的某个维度分到某组卡上」加上「因此在某处插入某种集合操作」。框架的差别在于这两件事由谁来写。

Megatron-LM:手写。 每种并行是一组显式的模块和一组进程组。张量并行是 ColumnParallelLinearRowParallelLinear 两个类,ffgg 是它们里面的自定义 autograd 函数;序列并行是这两个类的一个开关,把 all-reduce 换成 all-gather 和 reduce-scatter;流水线是一个调度器(1F1B、interleaved),段间用点对点通信;expert 并行是 MoE 层里的 dispatcher,负责 all-to-all 和 token 的排列;上下文并行是注意力算子的一个包装,做 KV 轮转。每一种并行有自己的进程组,tensor_model_parallel_sizepipeline_model_parallel_sizeexpert_model_parallel_sizecontext_parallel_size 四个参数加上剩下的卡数就是数据并行度,进程组按固定的顺序生成(默认 tp-cp-ep-dp-pp,tp 变化最快)。DeepSpeed 走的是同一条路,只是重心在数据并行这一侧:ZeRO 的三级分片、Ulysses 的序列并行、Offload,流水线和 MoE 也有自己的引擎。

PyTorch DTensor / FSDP2 / torchtitan:写放置,通信自动来。 这一路的核心抽象是 DeviceMeshDTensor。mesh 是一个多维的卡阵列,每个维度有名字;DTensor 是一个带「放置」(placement)的张量,放置有三种:Shard(k) 沿第 kk 维切到 mesh 的某一维上,Replicate() 每张卡一份,Partial() 每张卡有一个部分和、还没相加。六种并行在这个语言里是:

并行谁是 DTensor放置通信在哪里出现
DP激活沿 batch Shard(0),权重 Replicate反向的权重梯度是 Partial,转成 Replicate 要 all-reduce
FSDP权重沿第 0 维 Shard(0)前向用到时转 Replicate,all-gather;梯度 PartialShard(0),reduce-scatter
TP权重AA Shard(1)(列),BB Shard(0)(行)XAXA 的输出是 Shard(1),乘 BBPartial,转 Replicate 要 all-reduce
SP激活在 LayerNorm 处沿序列 Shard(1)进 TP 前转 Replicate 是 all-gather;出 TP 时 PartialShard(1) 是 reduce-scatter
CPQ,K,VQ, K, V沿序列 Shard注意力算子内部的 KV 轮转
EPexpert 权重沿 expert 维 Shard(0)token 到 expert 的路由是一次 all-to-all
PP不是张量的切分模块按层拆到不同 stage点对点,由 torch.distributed.pipelining 的调度器发出

表里第三列每一行都是一条「从一种放置转到另一种放置」的规则,这叫 redistribute,通信量就是那一行集合操作的量。张量并行在这里根本没有单独的实现:把 AA 的放置写成 Shard(1)BB 写成 Shard(0),普通的 matmul 在 DTensor 上执行时,第一次的输出自然是 Shard(1),第二次的输出自然是 Partial,需要完整结果时 DTensor 自己插入 all-reduce。parallelize_module 里的 ColwiseParallelRowwiseParallelSequenceParallel 只是给模块的权重和输入输出指定放置的快捷方式。FSDP2 的 fully_shard 同理:把模块的参数变成 Shard(0) 的 DTensor,前向到达这个模块时 all-gather,反向离开时 reduce-scatter。torchtitan 是这套东西的参考组合:先建一个五维 mesh(大致按 pp、dp_replicate、dp_shard、cp、tp 的顺序,tp 最内),然后逐个套:parallelize_module 做 TP 和 SP,context_parallel 做 CP,fully_shard 在 dp_shard 维上做 FSDP(如果 dp_replicate 大于 1 就是 HSDP:节点内分片、节点间复制),最后用 pipelining 把模块切成 stage。每一层并行只看 mesh 的自己那一维,互不知道对方存在。

这就是 SPMD(single program, multiple data)的含义:所有 rank 跑同一份 Python 程序,程序里没有 if rank == 0,每张卡看到的只是同一个 DTensor 的不同分片,通信是从放置的不匹配里推导出来的,不是手写的。

JAX / GSPMD:编译器传播。 JAX 走得更远。用户只给部分张量标注 sharding(NamedSharding(mesh, PartitionSpec('data', 'model')),或者在 jitin_shardings 里),XLA 的 SPMD 分区器把这些标注沿着计算图传播到每个中间结果,决定每个算子在每张卡上算哪一块,并在需要的地方插入集合操作。Megatron 的张量并行在这里是「把 AA 的第 1 维放在 model 轴上」这一句标注,其余全部由编译器推出,包括 all-reduce 的位置。Alpa(2022)更进一步,把 sharding 的选择和层的流水线切分都交给搜索。代价是流水线这种依赖执行顺序的并行在纯 sharding 的语言里不自然,通常要另外处理。

三种表达的差别是灵活性和自动化的取舍,切法本身没有差别。Megatron 的 RowParallelLinear 和 DTensor 的 Shard(0) 权重是同一件事;Megatron 的 fg 和 DTensor 的 redistribute 也是同一件事。会了一种,另外两种是查文档的问题。

怎么叠:一张 device mesh

真实训练里六种并行是叠着用的,每一种占 mesh 的一个维度,卡数是各维度大小的乘积。哪种并行放哪一维,由通信量决定:

  • TP 放最内层,它每层通信四次、每次 2bSd2\,bSd,只能跑在节点内的 NVLink 上,所以 T8T \le 8,占 mesh 里 rank 变化最快的那一维。
  • CP 紧挨着 TP,它的 KV 轮转也是每层都有,能放节点内最好。
  • PP 跨节点,每段边界只传一个 bSdbSd
  • DP 放最外层,把剩下的卡数用完,每步一次 2Ψ2\Psi,能藏在整个反向后面。
  • EP 通常和 DP 共用同一组卡:expert 按 EP 切,非 expert 部分在这组卡上按 DP 复制,Megatron 里 expert_model_parallel_size 必须整除 DP 的度。
  • ZeRO / FSDP 沿 DP 那一维切状态,不占新的维度。
节点 0,8 卡 NVLink 全互联rank 0(0, 0, 0)rank 1(0, 0, 1)rank 2(0, 0, 2)rank 3(0, 0, 3)rank 4(0, 1, 0)rank 5(0, 1, 1)rank 6(0, 1, 2)rank 7(0, 1, 3)TP 组TP 组DP 组DP 组DP 组DP 组节点 1,8 卡 NVLink 全互联rank 8(1, 0, 0)rank 9(1, 0, 1)rank 10(1, 0, 2)rank 11(1, 0, 3)rank 12(1, 1, 0)rank 13(1, 1, 1)rank 14(1, 1, 2)rank 15(1, 1, 3)TP 组TP 组DP 组DP 组DP 组DP 组PP 组跨节点mesh 维度 (pp, dp, tp) = (2, 2, 4);rank = tp + 4·dp + 8·pp;变化最快的维度放在带宽最高的链路上每张卡上:的层内权重,的层,的 batch;加 ZeRO 则优化器状态再按 dp 组切
16 张卡、两个节点上的一个三维 mesh:。全局 rank ,tp 变化最快,所以一个 TP 组正好是节点内 NVLink 相连的四张卡;DP 组是节点内同一列的两张卡;PP 组是两个节点上位置相同的卡,走节点间网络。Megatron 的 rank 顺序、PyTorch 的 init_device_mesh、JAX 的 Mesh 表达的都是这一张图。

一个真实的选择可以说明这些规则怎么互相让步。Kimi K3 的预训练用了 PP(带虚拟段)、EP、ZeRO-1 数据并行、Pipeline ZeRO-2 梯度分片和 CP,没有 TP:它每层的非 expert 参数只有几亿,用不着切,而 2.7T 的 routed expert 用 EP 切比用 TP 切自然得多。后果在 Kimi K3 的系统篇里展开。

一张表收尾

每层每个 micro-batch 的通信量按 BF16 字节算,Ψ\Psi 是参数量。

并行切的维度为什么成立集合操作通信量每卡参数每卡激活
DPbatch梯度对样本可加all-reduce,每步一次4Ψ4\Psi / 步不变,ZeRO 后 /N/N/N/N
TP矩阵的行、列分块矩阵乘,非线性逐元素all-reduce,每层四次16bSd16\,bSd / 层/T/T,LN 等除外sbd(10+24/T+)sbd(10 + 24/T + \dots)
SP序列(LN、dropout)逐 token 算子,AR = RS + AGall-gather + reduce-scatter同 TP同 TPsbd(34/T+)sbd(34/T + \dots)
PP函数复合点对点4bSd4\,bSd / 段边界/P/P1F1B 下 stage rrPrP-r
EPexpert排列是精确的all-to-all,每层四次8kbSd8k\,bSd / 层expert /R/R随收到的 token 数
CP序列(注意力内)在线 softmax 可分块合并点对点轮转前向 4bSd\approx 4\,bSd,与 CC 无关不变/C/C

放在哪条链路上,上一节的 mesh 图已经说了:TP、CP 节点内,PP 跨节点,DP 最外层,EP 与 DP 共组。

这一章接下来讲什么

按表里的行来,每篇把一格展开:

  1. 流水线并行:1F1B、虚拟段与气泡。为什么气泡是 (P1)/m(P-1)/m,为什么 rank 0 攒的激活最多,虚拟段用什么换什么。
  2. 张量并行与序列并行:ffgg 的 autograd 实现,all-reduce 和矩阵乘的重叠,TT 怎么选。
  3. 数据并行与 ZeRO:优化器状态、梯度、参数三级分片各省多少;FSDP 的单步时序、HSDP 的二维 mesh。
  4. expert 并行:all-to-all 的实现,负载不均衡的几种解法。
  5. 上下文并行:ring attention 的反向、zigzag、Ulysses,以及线性注意力的状态传递。

显存和推理服务的通用部分在「性能优化」章:显存账本推理服务的缓存与调度