
CANN ops-transformer 通算融合算子实战BatchMatMulReduceScatterAllToAll 原理、shape 约束与 aclnn 两段式调用【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本篇技术指南围绕 CANN ops-transformer 仓库中的mc2/batch_mat_mul_reduce_scatter_allto_all算子展开系统讲解这一通算融合算子如何把 BatchMatMul 计算与 ReduceScatter、AllToAll 集合通信融合在一条流水线中执行覆盖产品支持、参数与 shape 约束、aclnn 两段式接口调用以及源码级实现原理。读完本文你将掌握该算子在 Atlas A3 超节点内的适用场景、输入输出维度的数学关系以及如何基于两段式 aclnn 接口编写可运行的多卡调用样例。产品支持情况BatchMatMulReduceScatterAllToAll 对硬件的支持范围非常明确仅支持 Atlas A3 系列产品其余产品均不支持产品是否支持Ascend 950PR/Ascend 950DT×Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品×Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×这一点在算子定义文件 batch_mat_mul_reduce_scatter_allto_all_def.cpp 中也有印证该算子只注册了ascend910_93这一 AICore 配置对应 A3 平台。功能说明与计算流程BatchMatMulReduceScatterAllToAll 是通算融合算子它将 BatchMatMul 矩阵计算与 ReduceScatter、AllToAll 两类集合通信并行编排让计算和通信在超节点内重叠执行从而减少中间张量的落盘与额外的通信轮次。整体计算流程为BatchMatMul 计算 → 转置仅 yShardType 等于 0 时需要→ ReduceScatter 集合通信 → Add加 bias→ AllToAll 集合通信。其中 y 为最终输出形式化描述如下$$ temp1 BatchMatMul(x, weight) $$$$ temp2 ReduceScatter(temp1) $$$$ temp3 Add(temp2, bias) $$$$ y AllToAll(temp3) $$从语义上看该算子覆盖的是 MoEMixture of Experts类大模型在 EP专家并行 TPTensor 并行混合并行下的典型数据流BatchMatMul 在 EP 域内各 rank 上计算各自专家分片ReduceScatter 把同一专家在不同 rank 上的部分结果按 TP 域规约分片Add 叠加 bias最后由 AllToAll 在 EP 域内把结果重新分布到各 rank得到最终输出。参数说明算子输入输出与属性的完整定义如下与 README.md 及 aclnn 接口文档 保持一致参数名输入/输出描述数据类型数据格式x输入BatchMatMul 计算的左矩阵必须为 3 维。FLOAT16、BFLOAT16NDweight输入BatchMatMul 计算的右矩阵数据类型与 x 保持一致必须为 3 维。FLOAT16、BFLOAT16NDbiasOptional输入Add 计算的 bias需在 ReduceScatter 通信后执行 Add 操作。x 为 FLOAT16 时biasOptional 需为 FLOAT16x 为 BFLOAT16 时biasOptional 需为 FLOAT32。支持两维或三维支持传入空指针。FLOAT16、FLOAT32NDgroupEp输入专家并行的通信域名字符串长度需大于 0 且小于 128。STRINGNDgroupTp输入Tensor 并行的通信域名字符串长度需大于 0 且小于 128。STRINGNDepWorldSize输入EP 通信域 size支持 2、4、8、16、32。INT64NDtpWorldSize输入TP 通信域 size支持 2、4、8、16、32。INT64NDyShardType输入整型0 表示在 H 维度BatchMatMul 计算结果的第 2 维结果共 3 维维度索引依次为 0、1、2按 tp 进行 ReduceScatter1 表示在 C 维度BatchMatMul 计算结果的第 1 维按 tp 进行 ReduceScatter。INT64NDout输出Device 侧的 aclTensor为 batch_matmul 计算 reduce_scatter 计算 all_to_all 通信的结果数据类型与输入 x 保持一致必须为 3 维。FLOAT16、BFLOAT16ND需要特别留意biasOptional的两点特性一是它的数据类型跟随 x 的类型x 为 FLOAT16 时 bias 必须同为 FLOAT16x 为 BFLOAT16 时 bias 反而要用 FLOAT32因为 BF16 精度有限Add 前需要提升到 FP32 计算再转回二是它允许传入空指针即可以省略 bias 参与 Add 这一步。op_api 层在 aclnn_batch_matmul_reduce_scatter_all_to_all.cpp 的CheckDtypeValid中对上述 dtype 组合做了逐一校验。另外从算子定义看weight还支持转置场景定义中有一个transpose_weight属性默认 false当传入的 weight 为最后两维转置的视图时aclnn 接口会通过TransTensor在 Host 侧构造出转置后的元数据再下发给内部接口。约束说明由于集合通信及 BatchMatMul 计算所需输入输出 shape 需满足以下数学关系其中 epepWorldSizetptpWorldSize按 H 轴进行 ReduceScatter 场景即 yShardType 为 0x:(E/ep, ep*C, M/tp)weight:(E/ep, M/tp, H)biasOptional非空指针情况下三维时为(E/ep, 1, H/tp)两维时为(E/ep, H/tp)y:(E, C, H/tp)按 C 轴进行 ReduceScatter 场景即 yShardType 为 1x:(E/ep, ep*tp*C/tp, M/tp)weight:(E/ep, M/tp, H)biasOptional非空指针情况下三维时为(E/ep, 1, H)两维时为(E/ep, H)y:(E, C/tp, H)数据关系与取值范围说明例如 x.size(0) 等于 E/tp、y.size(0) 等于 E 时表示y.size(0) ep * x.size(0)且 y.size(0) 是 ep 的整数倍其他关系类似。E 的取值范围为 [2, 512]且 E 是 ep 的整数倍。H 的取值范围为 [1, 65535]当 yShardType 为 0 时H 是 tp 的整数倍。M/tp 的取值范围为 [1, 65535]。E/ep 的取值范围为 [1, 32]。ep、tp 均仅支持 2、4、8、16、32。groupEp 和 groupTp 名称不能相同。C 大于 0上限为算子 device 内存上限当 yShardType 为 1 时C 是 tp 的整数倍。通算融合算子不支持并发调用不同的通算融合算子也不支持并发调用。不支持跨超节点只支持超节点内。这些约束不仅写在文档中也落实在代码里op_api 层的CheckAttr校验 ep/tp 取值CheckTensorDimCommonShape、CheckTensorDimUniqueShape校验y_0 x_0 * ep、x_2 w_1等公共/非公共维度关系CheckShapeRange校验 E、H、M/tp、E/ep 的范围见 aclnn_batch_matmul_reduce_scatter_all_to_all.cppbmm_reduce_scatter_all_to_all_infershape.cpp 则在图编译阶段完成同样的 shape/dtype 检查并推导输出 shapeyShardType 为 0 时输出(E, C, H/tp)为 1 时输出(E, C/tp, H)。调用说明aclnn 两段式接口本算子的 aclnn 接口遵循 CANN 的两段式调用范式必须先调用aclnnBatchMatMulReduceScatterAlltoAllGetWorkspaceSize获取计算所需 workspace 大小以及包含算子计算流程的执行器再调用aclnnBatchMatMulReduceScatterAlltoAll执行计算。aclnnStatus aclnnBatchMatMulReduceScatterAlltoAllGetWorkspaceSize( const aclTensor* x, const aclTensor* weight, const aclTensor* biasOptional, const char* groupEp, const char* groupTp, int64_t epWorldSize, int64_t tpWorldSize, int64_t yShardType, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor)aclnnStatus aclnnBatchMatMulReduceScatterAlltoAll( void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream)第一段接口除计算 workspace 外还会完成全部入参校验主要错误码如下返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入的 x、weight、groupEp、groupTp 或 out 是空指针。ACLNN_ERR_PARAM_INVALID1610021. groupEp 或 groupTp 字符串长度不合法2. 输入不支持的数据类型3. 属性值不合法4. aclTensor 维度不合法5. aclTensor shape 不合法。第二段接口的参数含义workspace为 Device 侧申请的 workspace 内存地址workspaceSize由第一段接口返回executor为包含算子计算流程的执行器stream为执行任务的流。两段接口均返回 aclnnStatus 状态码。完整调用示例8 rankEP4TP2样例仓库在 examples/test_aclnn_batch_mat_mul_reduce_scatter_allto_all.cpp 提供了完整的可运行样例同时在 tests/ut/op_api/test_aclnn_batch_mat_mul_reduce_scatter_allto_all.cpp 保留了同结构的单测。核心流程如下1. 初始化环境创建各 rank 的 Context 与 Streamint ret aclInit(nullptr); // 为每个 rank 设置设备、创建 context 与 stream for (uint32_t rankId 0; rankId DEV_NUM; rankId) { aclrtSetDevice(rankId); aclrtCreateContext(context[rankId], rankId); aclrtCreateStream(stream[rankId]); }2. 通过 HcclCommInitAll 初始化 EP 与 TP 通信域EP4、TP2 时共 8 个 rank。EP 域按{0,2,4,6}、{1,3,5,7}划分TP 域按{0,1}、{2,3}、{4,5}、{6,7}划分// 初始化 ep 域 ep 4{0,2,4,6} {1,3,5,7} for (int i 0; i TP_WORLD_SIZE; i) { for (int j 0; j EP_WORLD_SIZE; j) { devicesEp[j i * EP_WORLD_SIZE] i j * TP_WORLD_SIZE; } HcclCommInitAll(EP_WORLD_SIZE, devicesEp[i * EP_WORLD_SIZE], commsEp[i * EP_WORLD_SIZE]); } // 初始化 tp 域 tp 2{0,1} {2,3} {4,5} {6,7} for (int i 0; i EP_WORLD_SIZE; i) { for (int j 0; j TP_WORLD_SIZE; j) { devicesTp[j i * TP_WORLD_SIZE] j i * TP_WORLD_SIZE; } HcclCommInitAll(TP_WORLD_SIZE, devicesTp[i * TP_WORLD_SIZE], commsTp[i * TP_WORLD_SIZE]); }3. 根据 yShardType 构造输入输出 shape样例选取E 4*EP、C 6*TP、H 2*TP、M 6*TP默认yShardType 1按 C 维 ReduceScatterif (xShardType 1) { xShape {E / EP_WORLD_SIZE, EP_WORLD_SIZE * TP_WORLD_SIZE * C / TP_WORLD_SIZE, M / TP_WORLD_SIZE}; weightShape {E / EP_WORLD_SIZE, M / TP_WORLD_SIZE, H}; biasShape {E / EP_WORLD_SIZE, 1, H}; yOutShape {E, C / TP_WORLD_SIZE, H}; } else if (xShardType 0) { xShape {E / EP_WORLD_SIZE, EP_WORLD_SIZE * C, M / TP_WORLD_SIZE}; weightShape {E / EP_WORLD_SIZE, M / TP_WORLD_SIZE, H}; biasShape {E / EP_WORLD_SIZE, 1, H / TP_WORLD_SIZE}; yOutShape {E, C, H / TP_WORLD_SIZE}; }注意通信域名称不能直接传 hccl_world_group而应通过HcclGetCommName从通信句柄上取回真实域名后再传入算子char hcomEpName[128] {0}; HcclGetCommName(args.hcclEpComm, hcomEpName); char hcomTpName[128] {0}; HcclGetCommName(args.hcclTpComm, hcomTpName);4. 两段式接口调用与资源回收// 调用第一阶段接口获取 workspace 大小与执行器 ret aclnnBatchMatMulReduceScatterAlltoAllGetWorkspaceSize( x, weight, bias, hcomEpName, hcomTpName, EP_WORLD_SIZE, TP_WORLD_SIZE, xShardType, yOut, workspaceSize, executor); // 按 workspaceSize 申请 device 内存 if (workspaceSize 0) { aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 调用第二阶段接口执行计算 ret aclnnBatchMatMulReduceScatterAlltoAll(workspaceAddr, workspaceSize, executor, args.stream); // 同步等待任务执行结束固定写法 ret aclrtSynchronizeStreamWithTimeout(args.stream, 10000);结束后依次aclDestroyTensor销毁张量、aclrtFree释放 device 内存、HcclCommDestroy销毁通信域、aclrtDestroyStream/aclrtResetDevice清理资源。5. 多线程并发执行由于算子涉及 8 个 rank 的集合通信样例为每个 rank 启动一个线程同时发起调用通信完成依赖所有 rank 同步参与最后 join 等待全部线程结束再aclFinalize()。这也是集合通信类算子的固定多进程/多线程编排方式——所有 rank 必须同时进入通信调用否则会造成通信挂死。源码级实现解析op_api 层的入参校验与 weight 转置处理在 aclnn_batch_matmul_reduce_scatter_all_to_all.cpp 中CheckParams按序完成空指针、通信域名字符串长度上限HCCL_GROUP_NAME_MAX、dtype、属性、维度、shape 范围六类校验。其中 shape 校验逻辑非常直观地映射了文档约束CheckTensorDimCommonShapey_0 x_0 * ep、x_0 w_0、x_2 w_1不转置时 x 最后一维等于 weight 第 1 维CheckTensorDimUniqueShapeyShardType 为 0 时w_2 y_2 * tp且x_1 y_1 * ep为 1 时w_2 y_2且x_1 y_1 * ep * tpCheckShapeRangeM/tp ∈ [1, 65535]、E/ep ∈ [1, 32]、E ∈ [2, 512]、H ∈ [1, 65535]。当检测到 weight 的 view 为最后两维转置时会调用TransTensor构造一个交换了第 1、2 维 shape 与 stride 的临时 aclTensor并将transposeWeighttrue传给内部接口aclnnInnerBatchMatMulReduceScatterAlltoAllGetWorkspaceSize。tiling通信域配置、Lite 模式与 workspace 估算tiling 主流程在 batch_matmul_reduce_scatter_all_to_all_tiling.cpp 中从属性中取出 ep/tp/yShard 等参数填充commonTilingEOverEp、COverTp、H、MOverTp、epGroupSize、tpGroupSize 等调用公式化 tilingReduceScatterAll2AllBMM/ReduceScatterAll2AllBMMShardH切分 local本 rank 负责的 EP 分片与 non-local其他 EP rank 的分片块SetHcclTiling为两种集合通信配置算法ReduceScatterlevel0:doublering环形算法、AlltoAlllevel0:fullmesh;level1:pairwise超节点内全互联 两级 pairwise并将 ep/tp 通信域名称、HCCL 数据类型写入 tiling 数据MC2SetWorkspace/MC2SetWorkspaceShard根据2*E*C*HOverTp E*C*H的通信缓冲规模叠加 16MB 冗余估算 workspace当yShardType 1且COverTp 640时启用 Lite 模式isLite此时输入 A 在 BMM 前需要先转置kernel 会额外用 workspace 承载转置缓冲。tiling 最终通过UpdateTilingKey把yShardFlag/isWeightTrans/isBias/isLite组合成 tiling key驱动 kernel 在运行期选择对应模板实例。kernelNon-Local 与 Local 两阶段流水kernel 主类在 batch_mat_mul_reduce_scatter_allto_all.h 中整体分为NonLocalCommunicationAndCal()与LocalCommunicationAndCal()两个阶段Non-Local 阶段对 ep 域内其他 rank 的专家分片执行 BatchMatMulAIC 核上的BmmNonLocalCal随后在 AIV 核上发起 ReduceScatterhcclReduceScatter.ReduceScatterfalse归约操作HCCL_REDUCE_SUM、等待结果、执行 AddTransposeAddTransposeBeforeAlltoAll其中 BF16 输入会先 Cast 成 FP32 做 Add 再转回最后下发 AllToAllVLocal 阶段处理本 rank 的专家分片BMM 计算BmmLocalCal时跳过其他 ep rank 的 A 分片if (j epRankId) continue;随后以InterHcclGroupSync与 Handle 机制让 ReduceScatter 与 AlltoAll 交替流水执行最后统一等待全部 AllToAll 完成WaitAllAlltoAll计算与通信通过SyncAll、HcclHandle 的 Wait/Commit 完成核间与通信间的同步Process()末尾在 AIV 核上Finalize两个 Hccl 句柄。从源码结构看这种先算 non-local 分片并立刻通信、再算 local 分片的编排正是为了让 BatchMatMul 计算与 ReduceScatter/AllToAll 通信在超节点内最大化重叠。该 kernel 实现位于 op_kernel/arch22/batch_mat_mul_reduce_scatter_allto_all.cpp模板参数由batch_mat_mul_reduce_scatter_allto_all_tiling_key.h与_tiling_struct.h定义。op_graphKFC 任务生成op_graph/bmm_reduce_scatter_all_to_all_gen_task.cpp 通过Mc2GenTaskOpsUtils::CommonKFCMc2CalcParamFunc将算子的计算参数交给 aicpu kfc serverreuse key 为kfc_stream统一处理并复用Mc2MoeGenTaskOpsUtils::Mc2MoeGenTaskCallback生成集合通信类任务与算子定义中this-MC2().HcclGroup({group_ep, group_tp})声明的通信域绑定逻辑一致。总结BatchMatMulReduceScatterAllToAll 是 ops-transformer 中面向 MoE 大模型 EPTP 混合并行场景的通算融合算子它把 BatchMatMul、ReduceScatter、Add、AllToAll 四级流水融合在超节点内通过非本地/本地两阶段编排让计算与通信重叠。使用时需要严格遵循其 shape 数学关系与取值约束E ∈ [2, 512]、H/M-tp ∈ [1, 65535]、ep/tp ∈ {2,4,8,16,32} 等并通过两段式 aclnn 接口完成 workspace 申请与计算下发。更多细节可进一步阅读 接口文档、示例代码 以及 算子定义 与 tiling 实现。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考