
ColossalAI 2D 张量并行基于 SUMMA 的二维模型切分原理与实现【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI本文围绕 ColossalAI 张量并行系列文档中的《2D 张量并行》展开系统讲解 1D 张量并行在激活值activations内存上的瓶颈、2D 张量并行如何基于 SUMMA 矩阵乘法算法将输入与权重同时按q × q网格切分以及其计算/内存/通信效率的定量分析。结合仓库中colossalai/legacy内实际的 2D 并行层实现与进程组初始化源码读者可以完整掌握 2D 张量并行的算法推导、处理器网格编排要求与代码级实现细节并了解它在新版本与 Shardformer 中的集成现状。本文是 ColossalAI 张量并行教程系列的一部分建议先阅读 1D 张量并行 以获得前置知识同系列还包括 2.5D 张量并行 与 3D 张量并行可在掌握 2D 方案后继续进阶。引言1D 张量并行的瓶颈与 2D 方案的动机在 1D 张量并行 中模型的权重矩阵只沿单一维度被切分如列并行A [A1 A2]或行并行因此参数parameters会被均匀分摊到各个处理器上内存占用约降为O(1/P)。但 1D 张量并行没有对输入激活值activations进行划分——每个处理器仍持有完整激活张量的一份副本。对于大规模模型而言激活值在前向传播与反向传播以及自动微分所需的中间结果中占据的内存同样非常可观这会成为显存瓶颈之一。为了平均分配计算与内存负荷ColossalAI 文档引入了基于SUMMAScalable Universal Matrix Multiplication Algorithm可扩展通用矩阵乘法算法的 2D 张量并行方案。其核心思想是把参与 GEMM 的两个矩阵——输入与权重——都组织成二维块结构从而让参数与激活值都获得二次方的均摊。该方法对应的学术出处是论文An Efficient 2D Method for Training Super-Large Deep Learning Models对应文档见 2D_tensor_parallel.md。SUMMA 算法以线性层 Y XA 为例的逐步推导以线性层Y XA为例说明 2D 张量并行的完整计算过程。给定处理器个数P q × q这是必要条件即总处理器数必须是某个整数q的平方处理器按q行q列排成二维逻辑网格。当q 2时共 4 个处理器把输入X与权重A都按q × q划分为 4 个块$$ \left[\begin{matrix} X_{00} X_{01} \ X_{10} X_{11} \end{matrix} \right] \text{~and~} \left[\begin{matrix} A_{00} A_{01} \ A_{10} A_{11} \end{matrix} \right]. $$注意这里X的二维划分同时覆盖了两个维度因此相比 1D 方案每个处理器持有的激活分块只有 1/q²这正是 2D 方案能削减激活值内存的根本原因。整个计算由q 步组成第 t 1 步X_{i0}在其所在行内被广播即第 0 列持有X_{i0}的处理器把数据广播给同行的其他处理器A_{0j}在其所在列内被广播即第 0 行持有A_{0j}的处理器把数据广播给同列的其他处理器。于是每个处理器(i, j)手上都有了对应的输入块与权重块$$ \left[\begin{matrix} X_{00},A_{00} X_{00},A_{01} \ X_{10},A_{00} X_{10},A_{01} \end{matrix} \right]. $$随后在每个处理器(i, j)上执行一次本地矩阵乘法X_{i0}×A_{0j}$$ \left[\begin{matrix} X_{00}A_{00} X_{00}A_{01} \ X_{10}A_{00} X_{10}A_{01} \end{matrix} \right] \quad (1). $$第 t 2 步同理X_{i1}在行内广播、A_{1j}在列内广播再做一次本地 GEMM$$ \left[\begin{matrix} X_{01}A_{10} X_{01}A_{11} \ X_{11}A_{10} X_{11}A_{11} \end{matrix} \right] \quad (2). $$将(1)与(2)的结果在本地累加即得到完整的矩阵乘积$$ Y XA \left[\begin{matrix} X_{00}A_{00}X_{01}A_{10} X_{00}A_{01}X_{01}A_{11} \ X_{10}A_{00}X_{11}A_{10} X_{10}A_{01}X_{11}A_{11} \end{matrix} \right]. $$可以看到整个过程把一个大 GEMM 拆成了q²份相互独立的本地 GEMM 片段并通过 q 轮行广播 列广播 本地乘累加完成全局归约——这正是 SUMMA 算法的精髓每轮只广播一条条带panel用计算掩盖通信用分布式累加替代集中式 All-Reduce 大矩阵。源码层面的 SUMMA 印证这段推导在仓库的 legacy 实现中有直接对应。Matmul_AB_2D定义于 colossalai/legacy/nn/layer/parallel_2d/_operation.py的forward中输入激活A即上文的X在row_parallel_modeParallelMode.PARALLEL_2D_ROW进程组内以dist.broadcast异步广播权重B即上文的权重矩阵在col_parallel_modeParallelMode.PARALLEL_2D_COL进程组内以dist.broadcast异步广播外层for i in range(summa_dim)恰好循环 q 步每步torch.addmm(C, A_list[cur], B_list[cur], outC)完成一次本地乘累加与文档中 q 步推导一一对应。实现上还使用了双缓冲环形队列A_list、B_list各分配 2 个缓冲让当前步的本地 GEMM 与下一步的广播通信重叠执行以隐藏通信时延# 摘自 colossalai/legacy/nn/layer/parallel_2d/_operation.py 的 Matmul_AB_2D.forward opa[0] dist.broadcast(A_list[0], srcsrc_a, grouprow_group, async_opTrue) opb[0] dist.broadcast(B_list[0], srcsrc_b, groupcol_group, async_opTrue) for i in range(summa_dim): if i ! summa_dim - 1: # 预取下一步要广播的分块与当前计算重叠 opa[1 - cur] dist.broadcast(A_list[1 - cur], srcsrc_a 1, grouprow_group, async_opTrue) opb[1 - cur] dist.broadcast(B_list[1 - cur], srcsrc_b summa_dim, groupcol_group, async_opTrue) opa[cur].wait() opb[cur].wait() torch.addmm(C, A_list[cur], B_list[cur], outC)backward中输入梯度与权重梯度分别经由Matmul_ABT_2D计算ABᵀ与Matmul_ATB_2D计算AᵀB两个额外的 SUMMA 风格 GEMM 原语求得——这与 1D 张量并行文档中反向需要聚合的说明一脉相承由于每个处理器只持有完整权重的 1/q²梯度计算同样需要跨进程组的归约与广播。这三个矩阵乘法原语Matmul_AB_2D/Matmul_ABT_2D/Matmul_ATB_2D在仓库中都有对应的算子级测试见 tests/test_legacy/test_layers/test_2d/checks_2d/check_operation_2d.py。效率分析计算、内存与通信成本给定P q × q个处理器基于**环形算法ring algorithm**的 2D 张量并行的前向与反向的理论成本如下表完整引自 2D_tensor_parallel.md计算内存 (参数)内存 (activations)通信 (带宽)通信 (时延)O(1/q²)O(1/q²)O(1/q²)O(6(q-1)/q)O(6(q-1))逐项解读计算与参数内存均为O(1/q²)权重矩阵被切成q × q块每个处理器只持有1/q²的参数并只承担1/q²的浮点运算。激活值内存也达到O(1/q²)这是 2D 相比 1D 最关键的优势——1D 方案中激活值内存为O(1)不切分2D 方案把输入也切成q × q块因此激活值的均摊同样是二次方的。带宽O(6(q-1)/q)与 1D 的O(2(P-1)/P)形态相似由于每个处理器在 q 轮 SUMMA 循环中都参与行/列广播通信量随切分粒度增长平缓时延O(6(q-1))随 q 线性增长q 轮串行广播意味着 q 越大、通信轮次越多这是 2D 方案在大规模处理器网格下的主要代价。处理器规模约束P q × q意味着可并行度必须取完全平方数例如 4q2、9q3、16q4。这与 1D 张量并行任意 P 个处理器的灵活性不同是 2D 张量并行在工程部署时最直接的约束。代码中这一断言体现在 2D 进程组初始化器里# colossalai/legacy/context/process_group_initializer/initializer_2d.py self.summa_dim int(math.sqrt(self.tensor_parallel_size)) assert self.tensor_parallel_size self.summa_dim**2, 2D summa dim should equal to tensor parallel size ^ 0.5代码实现2D 并行的分层结构与关键算子进程组编排二维逻辑网格2D 张量并行在 ColossalAI 中引入了两套正交的并行模式。在 colossalai/legacy/context/parallel_mode.py 中可以看到枚举定义# 2D parallel PARALLEL_2D_ROW 2d_row PARALLEL_2D_COL 2d_col其中PARALLEL_2D_ROW2D 行组沿网格行方向组织的进程组对应 SUMMA 中输入在行内广播PARALLEL_2D_COL2D 列组沿网格列方向组织的进程组对应 SUMMA 中权重在列内广播。两者的划分逻辑分别由 initializer_2d.py 中的Initializer_2D_Row与Initializer_2D_Col完成二者通过统一的入口Initializer_2D注册DIST_GROUP_INITIALIZER.register_module在分布式初始化时依据tensor_parallel_size自动算出summa_dim √P并创建上述两组 NCCL 进程组同时为 CPU 通信创建 gloo 组。层内部还会读取环境变量中的SUMMA_DIM来获得切分维度相关校验在 parallel_2d/_utils.py 的get_summa_dim_from_env与assert_summa_initialization中完成。2D 并行层族colossalai/legacy/nn/layer/parallel_2d/layers.py 提供了一套完整的*2D层族均通过LAYERS.register_module注册涵盖构造典型 Transformer/CNN 所需的各类算子Linear2D2D 并行线性层。每个处理器持有一块形状为[in_features/q, out_features/q]的权重对应上文A_{ij}偏置则进一步切到out_features/q²粒度并通过add_bias_2d内部先all_gather偏置、再按 SUMMA 格局广播完成加偏置详见 parallel_2d/_operation.py。其forward直接调用上文分析的Matmul_AB_2D。LayerNorm2D2D 并行 LayerNorm。由于归一化维度的特征被分布在行组内均值与方差统计量需要先在PARALLEL_2D_ROW组上做all_reduce再通过layernorm_2d完成分布式归一化最后用 2D 加偏置原语应用weight/bias。Embedding2D/VocabParallelEmbedding2D2D 并行的词嵌入层其中 vocab-parallel 版本把词表按1/q切分、嵌入维按1/q切分并结合reduce_scatter输出用于语言模型的 embedding 与输出头。PatchEmbedding2D面向 ViT 等视觉模型的 2D 并行 Patch 嵌入层在列组上all_gather卷积核后再执行F.conv2d并同样切分 cls token 与位置编码。Classifier2D2D 并行分类输出层in_features按1/q²切分。这些层都通过set_tensor_parallel_attribute_by_partition(..., self.summa_dim**2)标记被切分了q²份的属性供后续 checkpoint 的切分/合并逻辑识别。相应地层级的加载与保存借助partition_tensor_parallel_state_dict/gather_tensor_parallel_state_dict先在行组、后在列组两个方向分别切分/聚合见 layers.py 中的_load_from_global_state_dict与_save_to_global_state_dict。关键通信原语parallel_2d/_operation.py 中还实现了若干面向 2D 格局的封装原语均可自动微分、支持 AMPcustom_fwd/custom_bwdall_gather_tensor_2d沿指定维度在对应并行模式组上做 all-gather反向为 reduce-scatterreduce_scatter_tensor_2dreduce-scatter反向为 all-gather如VocabParallelEmbedding2D的输出归约reduce_tensor_2d/reduce_by_batch_2d跨模型并行区域的 all-reduce用于统计量如 batch 维规约的 loss在 2D 网格上的同步split_batch_2d在PARALLEL_2D_COL维度上按批次切分输入供 2D 网格按 batch 分摊输入数据。使用方式与现状Shardformer 集成、旧版 API 与测试关于集成现状以当前仓库文档为准文档明确指出原文见 2D_tensor_parallel.md 的使用一节ColossalAI 的最新版本还暂不支持 2D 张量并行但 2D 张量并行的功能会在未来的版本被集成入Shardformer中。也就是说在**新版新 API 时代**中面向 HuggingFace 主流 Transformer 模型的自动并行切分由 Shardformer 承担而 Shardformer 当前聚焦于 1D 张量并行以及流水线并行、序列并行、FlashAttention、JIT 融合算子等优化2D 张量并行尚未被集成进 Shardformer因此本文所述算法主要面向理解原理以及作为 legacy 代码的参考实现若希望直接体验 2D 张量并行的运行效果可参考旧版本 ColossalAI 的用法见下文legacy 实现与测试并留意未来版本 Shardformer 对 2D 支持的更新。legacy 实现与测试如何验证 2D 并行的正确性虽然 2D 张量并行未被纳入新版的 Shardformer本仓库仍在colossalai/legacy/目录保留了完整的 2D 张量并行实现链上文引用的layers.py、_operation.py、进程组初始化器等均在此目录并有配套的端到端测试可以直接运行验证。在旧版legacy分布式配置中通过launch并指定张量并行模式即可启用 2D 并行。仓库测试给出的标准配置是 4 卡q2场景见 tests/test_legacy/test_layers/test_2d/test_2d.pyCONFIG dict( paralleldict(pipelinedict(size1), tensordict(size4, mode2d)), ) def check_layer_and_operation(rank, world_size, port): disable_existing_loggers() launch(rankrank, world_sizeworld_size, hostlocalhost, portport, backendnccl) torch.backends.cuda.matmul.allow_tf32 False torch.backends.cudnn.allow_tf32 False torch.backends.cudnn.deterministic True check_layer() gpc.destroy() def test_2d(): spawn(check_layer_and_operation, 4)关键点解读tensor.size必须为完全平方数tensor.mode2d触发 2D 进程组初始化4 卡时q 2。测试中显式关闭 TF32、开启确定性 CUDA 后端以保证分布式 GEMM 与单卡基准逐位可比。check_layer()依次校验Linear2D、LayerNorm2D、Embedding2D、PatchEmbedding2D、带/不带给定权重的Classifier2D、vocab-parallel 变体以及 2D loss 的数值正确性前向输出与梯度对齐全局参数其实现位于 checks_2d/check_layer_2d.py算子级测试则位于 checks_2d/check_operation_2d.py。注意上述 legacy API 属于 ColossalAI 旧版本接口在当前仓库中位于colossalai/legacy下实际使用时应以你所安装的 ColossalAI 版本所支持的 API 为准。该目录内的实现与测试主要是为了让 2D 张量并行的算法细节在源码层面可被研读与验证。使用 2D 并行时的实用建议综合算法特性与实现约束在实际选型中可参考以下几点处理器数必须为q²规划资源时优先考虑 4、9、16、25 等完全平方数的卡数否则无法启用 2D 并行分块尺寸需要整除in_features、out_features需要能被qLinear2D的权重维乃至q²偏置维、LayerNorm2D的归一化维整除batch 维需要能被列组大小整除否则会触发断言失败如_operation.py中reduce_scatter_tensor_2d对dim_size % world_size的检查激活值敏感场景收益最大当模型中间激活张量巨大、1D 张量并行受限于激活内存时2D 方案把激活内存额外降至1/q²是值得关注的方向权衡通信时延2D 的时延开销随q线性增长表格中O(6(q-1))在q过大、网络带宽与拓扑不佳时收益可能被通信抵消需要结合实际集群做基准测试。总结核心收益相比只切参数的 1D 张量并行基于 SUMMA 的 2D 张量并行把输入与权重同时按q × q分块使计算、参数内存与激活值内存均摊到O(1/q²)算法本质在q轮迭代中通过行内广播输入块 列内广播权重块 本地 GEMM 累加复现完整矩阵乘法代码上的落点即 Matmul_AB_2D 中带双缓冲重叠的 q 步循环工程约束处理器数必须是完全平方数q²各类张量维度需要满足整除约束现状与展望当前新版 ColossalAI 中 2D 张量并行尚未并入 Shardformer本文结合colossalai/legacy/nn/layer/parallel_2d/与 test_2d.py 给出了原理、实现与数值验证的完整闭环可作为理解更高维方案2.5D、3D的基石。【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考