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) $$
从语义上看,该算子覆盖的是 MoE(Mixture of Experts)类大模型在 EP(专家并行)+ TP(Tensor 并行)混合并行下的典型数据流:BatchMatMul 在 EP 域内各 rank 上计算各自专家分片,ReduceScatter 把同一专家在不同 rank 上的部分结果按 TP 域规约分片,Add 叠加 bias,最后由 AllToAll 在 EP 域内把结果重新分布到各 rank,得到最终输出。
参数说明
算子输入输出与属性的完整定义如下(与 README.md 及 aclnn 接口文档 保持一致):
| 参数名 | 输入/输出 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| x | 输入 | BatchMatMul 计算的左矩阵,必须为 3 维。 | FLOAT16、BFLOAT16 | ND |
| weight | 输入 | BatchMatMul 计算的右矩阵,数据类型与 x 保持一致,必须为 3 维。 | FLOAT16、BFLOAT16 | ND |
| biasOptional | 输入 | Add 计算的 bias,需在 ReduceScatter 通信后执行 Add 操作。x 为 FLOAT16 时,biasOptional 需为 FLOAT16;x 为 BFLOAT16 时,biasOptional 需为 FLOAT32。支持两维或三维,支持传入空指针。 | FLOAT16、FLOAT32 | ND |
| groupEp | 输入 | 专家并行的通信域名,字符串长度需大于 0 且小于 128。 | STRING | ND |
| groupTp | 输入 | Tensor 并行的通信域名,字符串长度需大于 0 且小于 128。 | STRING | ND |
| epWorldSize | 输入 | EP 通信域 size,支持 2、4、8、16、32。 | INT64 | ND |
| tpWorldSize | 输入 | TP 通信域 size,支持 2、4、8、16、32。 | INT64 | ND |
| yShardType | 输入 | 整型,0 表示在 H 维度(BatchMatMul 计算结果的第 2 维,结果共 3 维,维度索引依次为 0、1、2)按 tp 进行 ReduceScatter;1 表示在 C 维度(BatchMatMul 计算结果的第 1 维)按 tp 进行 ReduceScatter。 | INT64 | ND |
| out | 输出 | Device 侧的 aclTensor,为 batch_matmul 计算 + reduce_scatter 计算 + all_to_all 通信的结果,数据类型与输入 x 保持一致,必须为 3 维。 | FLOAT16、BFLOAT16 | ND |
需要特别留意biasOptional的两点特性:一是它的数据类型跟随 x 的类型,x 为 FLOAT16 时 bias 必须同为 FLOAT16,x 为 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 需满足以下数学关系(其中 ep=epWorldSize,tp=tpWorldSize):
按 H 轴进行 ReduceScatter 场景,即 yShardType 为 0:
- x:
(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 为 1:
- x:
(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.cpp);bmm_reduce_scatter_all_to_all_infershape.cpp 则在图编译阶段完成同样的 shape/dtype 检查并推导输出 shape(yShardType 为 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_NULLPTR | 161001 | 传入的 x、weight、groupEp、groupTp 或 out 是空指针。 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 1. groupEp 或 groupTp 字符串长度不合法;2. 输入不支持的数据类型;3. 属性值不合法;4. aclTensor 维度不合法;5. aclTensor shape 不合法。 |
第二段接口的参数含义:workspace为 Device 侧申请的 workspace 内存地址,workspaceSize由第一段接口返回,executor为包含算子计算流程的执行器,stream为执行任务的流。两段接口均返回 aclnnStatus 状态码。
完整调用示例:8 rank(EP=4,TP=2)样例
仓库在 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 与 Stream
int 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 通信域
EP=4、TP=2 时共 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 维 ReduceScatter):
if (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 校验逻辑非常直观地映射了文档约束:
CheckTensorDimCommonShape:y_0 == x_0 * ep、x_0 == w_0、x_2 == w_1(不转置时 x 最后一维等于 weight 第 1 维);CheckTensorDimUniqueShape:yShardType 为 0 时w_2 == y_2 * tp且x_1 == y_1 * ep;为 1 时w_2 == y_2且x_1 == y_1 * ep * tp;CheckShapeRange:M/tp ∈ [1, 65535]、E/ep ∈ [1, 32]、E ∈ [2, 512]、H ∈ [1, 65535]。
当检测到 weight 的 view 为"最后两维转置"时,会调用TransTensor构造一个交换了第 1、2 维 shape 与 stride 的临时 aclTensor,并将transposeWeight=true传给内部接口aclnnInnerBatchMatMulReduceScatterAlltoAllGetWorkspaceSize。
tiling:通信域配置、Lite 模式与 workspace 估算
tiling 主流程在 batch_matmul_reduce_scatter_all_to_all_tiling.cpp 中:
- 从属性中取出 ep/tp/yShard 等参数,填充
commonTiling(EOverEp、COverTp、H、MOverTp、epGroupSize、tpGroupSize 等); - 调用公式化 tiling(
ReduceScatterAll2AllBMM/ReduceScatterAll2AllBMMShardH)切分 local(本 rank 负责的 EP 分片)与 non-local(其他 EP rank 的分片)块; SetHcclTiling为两种集合通信配置算法:ReduceScatter=level0:doublering(环形算法)、AlltoAll=level0: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 在运行期选择对应模板实例。
kernel:Non-Local 与 Local 两阶段流水
kernel 主类在 batch_mat_mul_reduce_scatter_allto_all.h 中,整体分为NonLocalCommunicationAndCal()与LocalCommunicationAndCal()两个阶段:
- Non-Local 阶段:对 ep 域内其他 rank 的专家分片执行 BatchMatMul(AIC 核上的
BmmNonLocalCal),随后在 AIV 核上发起 ReduceScatter(hcclReduceScatter.ReduceScatter<false>,归约操作HCCL_REDUCE_SUM)、等待结果、执行 Add+Transpose(AddTransposeBeforeAlltoAll,其中 BF16 输入会先 Cast 成 FP32 做 Add 再转回),最后下发 AllToAllV; - Local 阶段:处理本 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_graph:KFC 任务生成
op_graph/bmm_reduce_scatter_all_to_all_gen_task.cpp 通过Mc2GenTaskOpsUtils::CommonKFCMc2CalcParamFunc将算子的计算参数交给 "aicpu kfc server"(reuse key 为kfc_stream)统一处理,并复用Mc2MoeGenTaskOpsUtils::Mc2MoeGenTaskCallback生成集合通信类任务,与算子定义中this->MC2().HcclGroup({"group_ep", "group_tp"})声明的通信域绑定逻辑一致。
总结
BatchMatMulReduceScatterAllToAll 是 ops-transformer 中面向 MoE 大模型 EP+TP 混合并行场景的通算融合算子,它把 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),仅供参考