ColossalAI 2.5D 张量并行深度解析:原理、切分语义、通信成本与工程实现
【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI
2.5D 张量并行(2.5D Tensor Parallelism)是 ColossalAI 在张量并行系列中承上启下的并行策略:它在 2D 张量并行(SUMMA)的基础上额外引入一层"深度(depth)"维度,用更多处理器把前向/反向中的通信量进一步压低。本文以线性层 $Y=XA$ 为例,完整推导 2.5D 的切分与计算过程,给出其理论计算/内存/通信复杂度,并结合本仓库legacy模块中的真实实现剖析进程组划分、SUMMA 通信循环与可用的并行层 API,帮助读者理清 2.5D 张量并行"为什么划算、怎么切、仓库里现在能用什么"。
本主题对应的原始资料为仓库文档 2p5D_tensor_parallel.md,它是 1D 张量并行 与 2D 张量并行 的后续内容,建议按 1D → 2D → 2.5D 的顺序阅读。
背景:为什么需要 2.5D
在进入数学细节之前,先明确 2.5D 要解决的问题:
- 1D 张量并行(Megatron-LM 风格)只把权重矩阵按列/行切开,激活值(activations)不做划分、在每个处理器上各存一份。对超大模型而言,激活值同样会消耗海量内存,且通信规模随处理器数线性增长(见 1D_tensor_parallel.md)。
- 2D 张量并行在 SUMMA(可扩展通用矩阵乘法)算法基础上,把输入和权重都切成 $q\times q$ 块,使计算与内存负载更均匀(见 2D_tensor_parallel.md),但它需要更多的通信步与更大的通信量。
- 2.5D 张量并行是在论文《2.5-dimensional distributed model training》中提出的,基于2.5D SUMMA:通过在网格上额外叠加一个"深度"维度来使用更多设备,从而在不降低内存收益的前提下显著压低通信成本。
一句话概括本文档的核心主张:相比 1D 降低了内存成本,相比 2D 用多出来的设备换取了更少的通信。
线性层 $Y=XA$ 的 2.5D 切分推演
设目标计算为一个线性层(GEMM):
$$ Y = XA $$
其中 $X$ 为输入(激活)、$A$ 为权重。为便于与 2D 对比,此处沿用文档中的符号约定。
处理器网格与"必要条件"
给定 $P = q \times q \times d$ 个处理器。这里有两个相互独立的结构参数:
- $q$:SUMMA 的"行数/列数"(tesseract 维度),权重和输入按 $q\times q$ 的块结构展开;
- $d$:2.5D 特有的深度维度(depth),即每个处理器在深度方向上属于 $d$ 层之一。
以 $q=d=2$ 为例,一共需要 $P=2\times2\times2=8$ 个处理器。源码中把 $q$ 记作tesseract_dim、把 $d$ 记作tesseract_dep(tesseract 即"立方体"之意),并显式校验 $q^2\times d$ 与张量并行总规模的一致性,见 initializer_2p5d.py 中对tensor_parallel_size == tesseract_dim**2 * tesseract_dep的断言。
输入 $X$ 的切分:$d\times q$ 行、$q$ 列
把输入 $X$ 划分为 $d\times q$ 行和 $q$ 列:
$$ \left[\begin{matrix} X_{00} & X_{01} \ X_{10} & X_{11} \ X_{20} & X_{21} \ X_{30} & X_{31}\end{matrix} \right] $$
它可以被重塑(reshape)为 $d$ 层,每层是一个 $q\times q$ 的块矩阵:
$$ \left[\begin{matrix} X_{00} & X_{01} \ X_{10} & X_{11} \end{matrix} \right] \text{~and~}\left[\begin{matrix} X_{20} & X_{21} \ X_{30} & X_{31} \end{matrix} \right]. $$
$X$ 的每个分片被交给 (row, col, dep) 三个维度对应位置上的处理器:前 $d$ 组行切片分别归属 $d$ 个深度层,同一深度层内再按 $q\times q$ 布局。这相当于在 2D 的二维网格之上,再沿深度方向复制了 $d$ 份输入分片。
权重 $A$ 的切分:$q\times q$ 块
权重 $A$ 被划分为 $q\times q$ 块:
$$ \left[\begin{matrix} A_{00} & A_{01} \ A_{10} & A_{11} \end{matrix} \right]. $$
注意这里的要点是:权重不做深度切分——每一层深度上的处理器看到的是同一份 $A$。因此参数在每个深度层都会被完整保留一份(参数冗余因子为 $d$),这与后面效率表中的"内存(参数)为 $O(1/q^2)$、不含 $1/d$"是一致的。
逐层执行 SUMMA
对 $X$ 的每一层,用 SUMMA 算法将 $X$ 与 $A$ 相乘。第一层得到:
$$ \left[\begin{matrix} Y_{00}=X_{00}A_{00}+X_{01}A_{10} & Y_{01}=X_{00}A_{01}+X_{01}A_{11} \ Y_{10}=X_{10}A_{00}+X_{11}A_{10} & Y_{11}=X_{10}A_{01}+X_{11}A_{11} \end{matrix} \right] $$
第二层得到:
$$ \left[\begin{matrix} Y_{20}=X_{20}A_{00}+X_{21}A_{10} & Y_{21}=X_{20}A_{01}+X_{21}A_{11} \ Y_{30}=X_{30}A_{00}+X_{31}A_{10} & Y_{31}=X_{30}A_{01}+X_{31}A_{11} \end{matrix} \right]. $$
从公式可以看出 2.5D SUMMA 与 2D SUMMA 的差异:每个深度层只需要对自己那 $q$ 行输入分片做一次完整的 $q$ 步 SUMMA。切分的输入行数更多($dq$ 行 vs $q$ 行),因此每一片输入更小;在固定 $q$ 的前提下,$d$ 越大,每个处理器承担的 GEMM 越小、且每个广播消息的体积越小。
后向传播
$Y=XA$ 的后向需要两类梯度:对权重 $A$ 的梯度 $\dot{A}=X^T\dot{Y}$,以及对输入 $X$ 的梯度 $\dot{X}=\dot{Y}A^T$。二者都需要跨处理器聚合分片结果。在仓库的算子实现中,后向通过reduce_scatter沿列并行组把权重梯度归约回各分片、对输入梯度的处理则依赖于对应的转置乘法算子(见下文源码分析),实现了与 SUMMA 前向对称的通信模式。
效率分析:计算、内存与通信复杂度
给定 $P=q \times q \times d$ 个处理器,基于环形(ring)算法的 2.5D 张量并行,其前向和后向的理论计算成本、内存成本与通信成本如下(原文档效率表):
| 计算 | 内存 (参数) | 内存 (activations) | 通信 (带宽) | 通信 (时延) |
|---|---|---|---|---|
| $O(1/dq^2)$ | $O(1/q^2)$ | $O(1/dq^2)$ | $\small O(3(q-1)(d+1)/dq)$ | $O(6(q-1))$ |
要点解读:
- 计算与激活内存都按 $1/(dq^2)$ 摊薄,即同时受益于 $q^2$(行列二维切分)与 $d$(深度复制);
- 参数内存只按 $1/q^2$ 摊薄,这是因为权重在每个深度层被复制 $d$ 份;
- 带宽型通信约为 $O(3(q-1)(d+1)/dq)$,随 $d$ 增大而趋近一个更小的常数(相比 2D 的 $O(6(q-1)/q)$,在带宽常数上明显下降);
- 时延型通信为 $O(6(q-1))$,与深度 $d$ 无关,体现为 SUMMA 的 $q$ 步结构。
为便于对照,把另外两档张量并行的同类指标一并列出(数据分别源自仓库文档 1D_tensor_parallel.md 与 2D_tensor_parallel.md):
| 策略 | 计算 | 内存(参数) | 内存(activations) | 通信(带宽) | 通信(时延) |
|---|---|---|---|---|---|
| 1D($P$ 个处理器) | $O(1/P)$ | $O(1/P)$ | $O(1)$ | $O(2(P-1)/P)$ | $O(2(P-1))$ |
| 2D($P=q^2$) | $O(1/q^2)$ | $O(1/q^2)$ | $O(1/q^2)$ | $O(6(q-1)/q)$ | $O(6(q-1))$ |
| 2.5D($P=q^2d$) | $O(1/dq^2)$ | $O(1/q^2)$ | $O(1/dq^2)$ | $O(3(q-1)(d+1)/dq)$ | $O(6(q-1))$ |
可以看到:2.5D 激活内存是三者中最低的一档,同时带宽通信常数随 $d$ 进一步收缩,代价是参数在每个深度层冗余 $d$ 份。因此 2.5D 适合激活规模极大、且带宽相比时延更"贵"的场景。
源码级剖析:legacy 模块中的 2.5D 实现
尽管新主线的Shardformer/booster尚未接入 2.5D,本仓库仍在colossalai/legacy/中保留了完整的 2.5D 工程实现,是研究其切分语义与通信细节的一手资料。
进程组划分:四组并行模式
2.5D 在进程组层面被拆成四组并行模式(定义于 parallel_mode.py):
PARALLEL_2P5D_ROW = "2p5d_row"—— 行并行组(同一深度层内、同一列上的处理器组成);PARALLEL_2P5D_COL = "2p5d_col"—— 列并行组(同一深度层内、同一行上的处理器组成);PARALLEL_2P5D_DEP = "2p5d_dep"—— 深度并行组(不同深度层上位置相同的处理器组成);PARALLEL_2P5D_XZ—— 行 × 深度组合组(用于需要同时跨行与跨深度通信的场合)。
进程组初始化逻辑见 initializer_2p5d.py:它根据tesseract_dim($q$)与tesseract_dep($d$)逐个遍历全局 rank 构造上述分组,并用环境变量tesseract_dim/tesseract_dep(tensor_parallel_env)对初始化参数做一致性校验。
2.5D 层族:一张现成的并行算子清单
2.5D 相关的层实现集中在 colossalai/legacy/nn/layer/parallel_2p5d/ 目录:
Linear2p5D:2.5D 线性层。构造时读取 row/col/dep 三维的本地 rank 与tesseract_dim,把输入维与输出维各除以 $q$(in_features // q、out_features // q),本地权重形状为 $[k/q,; h/q]$、偏置形状为 $[h/q]$;权重按 $q^2$ 分片打上张量并行属性标记。完整实现见 layers.py 中的 Linear2p5D。LayerNorm2p5D:在行并行组内先做一次all_reduce求局部和以得到全局均值/方差,再用 2.5D 专用算子完成归一化与缩放。Embedding2p5D/VocabParallelEmbedding2p5D:前者按隐藏维分片、前向时沿列组all_gather恢复完整词表;后者按词表维分片(VocabParallel),掩码掉非本分片的 token 后做reduce_scatter归约。Classifier2p5D:输出分类头,前向用all_gather+all_reduce在行/列组上聚合结果。PatchEmbedding2p5D:面向 ViT 类模型的 2D Patch Embedding 的 2.5D 版本(对 batch 做 2.5D 切分、沿列组 gather 权重)。
除层外,仓库还配套了 2.5D 的损失与指标实现:loss_2p5d.py 与 accuracy_2p5d.py。这些层均注册进LAYERS注册表,可通过字符串名称在模型构建中复用;模型权重在保存/加载时也会经由partition_tensor_parallel_state_dict/gather_tensor_parallel_state_dict完成分片与聚合,并先在深度组内广播保证每个深度层持有完整权重副本。
通信内核:SUMMA 循环与双缓冲流水
核心矩阵乘算子Matmul_AB_2p5D位于 _operation.py。它的前向结构精确对应了本文第三节的推演:
- 把输入 $X$(本地分片 $A$)与权重(本地分片 $B$)摊平成二维;
- 使用**两块环形缓冲(double buffer)**与异步
broadcast:第 $i$ 步沿行并行组广播输入分片、沿列并行组广播权重分片,源 rank 按tesseract_dim步长递进; - 等待通信完成后用
torch.addmm把本地乘积累加到结果 $C$ 上; - 循环 $q$ 次后得到完整乘积并 reshape 回输出形状。
关键点在于第 2 步的异步广播流水:每轮迭代会提前发起下一步的广播(async_op=True),与当前步的矩阵乘重叠执行,从而掩盖部分通信时延——代码注释中"2 is enough for all cases"正是指两块缓冲足以覆盖任意 $q$ 的流水。算子整体用@custom_fwd(cast_inputs=torch.float16)支持混合精度自动类型转换。
同文件还提供了配套的Matmul_ABT_2p5D(用于后向的转置乘法)、add_bias_2p5d、layernorm_2p5d、all_gather_tensor_2p5d、reduce_scatter_tensor_2p5d、split_batch_2p5d、reduce_by_batch_2p5d等 2.5D 专属通信算子(见 _operation.py 顶部类清单),这些算子封装了all_gather/all_reduce/reduce_scatter与对应反向梯度规约,供上面各层组合调用。
使用状态与工程注意点
当前主线:暂未接入 Shardformer
仓库文档明确说明:ColossalAI 最新版本暂不支持 2.5D 张量并行,该功能预期在未来的版本中集成进Shardformer。Shardformer是 ColossalAI 当前的模型并行引擎(自动把 HuggingFace 风格模型按策略切分并插入通信),其原理与用法详见 shardformer.md 与 shardformer 示例代码。也就是说,如果你使用的是新主线 API(booster+Shardformer),目前并不会看到 2.5D 配置项。
老版本用户 / 研究参考:legacy 实现仍然在仓库中
对于老版本 API 的使用方式,可以继续利用本仓库保留下来的实现:
- 并行模式字符串:
2p5d_row、2p5d_col、2p5d_dep、2p5d_xz; - 进程组初始化入口:initializer_2p5d.py,要求张量并行规模严格等于
tesseract_dim**2 * tesseract_dep; - 可直接复用的层:
Linear2p5D、LayerNorm2p5D、Embedding2p5D、VocabParallelEmbedding2p5D、Classifier2p5D、PatchEmbedding2p5D(位于 parallel_2p5d/layers.py); - 配套损失与指标:loss_2p5d.py、accuracy_2p5d.py;
- 完整的 2.5D 训练示例曾维护于官方示例仓库
hpcaitech/ColossalAI-Examples的features/tensor_parallel目录(历史版本),需要者可在该组织下检索对应 README;本仓库内可直接对照的 Shardformer 风格示例见 colossalai/shardformer/examples。
工程约束清单
从文档声明与源码结构可以归纳出使用 2.5D 时需要注意的硬性前提:
- 处理器数量必须满足 $P=q^2\times d$,即 $q$ 和 $d$ 需能从总张量并行规模中分解出来;
- 线性层的
in_features/out_features需能被 $q$ 整除(源码按divide(in_features, tesseract_dim)计算每片尺寸);归一化/词表维度同理; - 参数在每个深度层保存一份副本,模型体积较大的情况下应权衡 $d$ 的选择;
- 若在新 API 下使用,需等待 2.5D 被集成进
Shardformer,当前主线并不包含该策略的分发入口。
小结
2.5D 张量并行在 2D SUMMA 之上新增深度维 $d$,使计算与激活内存按 $O(1/dq^2)$ 摊薄、带宽通信随 $d$ 收缩为 $O(3(q-1)(d+1)/dq)$,同时付出参数冗余 $d$ 份的代价。仓库侧,新主线Shardformer尚未接入该策略,但 colossalai/legacy 中保留了一套从进程组(2p5d_row/col/dep)、SUMMA 双缓冲通信算子(Matmul_AB_2p5D)到可组合层(Linear2p5D等)的完整参考实现——对研究多维张量并行原理、或维护老版本训练脚本的读者而言,这套代码本身就是最直接的"活文档"。
【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考