news 2026/9/18 9:37:00

CANN ops-transformer 通算融合算子实战:BatchMatMulReduceScatterAllToAll 原理、shape 约束与 aclnn 两段式调用

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CANN ops-transformer 通算融合算子实战:BatchMatMulReduceScatterAllToAll 原理、shape 约束与 aclnn 两段式调用

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、BFLOAT16ND
weight输入BatchMatMul 计算的右矩阵,数据类型与 x 保持一致,必须为 3 维。FLOAT16、BFLOAT16ND
biasOptional输入Add 计算的 bias,需在 ReduceScatter 通信后执行 Add 操作。x 为 FLOAT16 时,biasOptional 需为 FLOAT16;x 为 BFLOAT16 时,biasOptional 需为 FLOAT32。支持两维或三维,支持传入空指针。FLOAT16、FLOAT32ND
groupEp输入专家并行的通信域名,字符串长度需大于 0 且小于 128。STRINGND
groupTp输入Tensor 并行的通信域名,字符串长度需大于 0 且小于 128。STRINGND
epWorldSize输入EP 通信域 size,支持 2、4、8、16、32。INT64ND
tpWorldSize输入TP 通信域 size,支持 2、4、8、16、32。INT64ND
yShardType输入整型,0 表示在 H 维度(BatchMatMul 计算结果的第 2 维,结果共 3 维,维度索引依次为 0、1、2)按 tp 进行 ReduceScatter;1 表示在 C 维度(BatchMatMul 计算结果的第 1 维)按 tp 进行 ReduceScatter。INT64ND
out输出Device 侧的 aclTensor,为 batch_matmul 计算 + reduce_scatter 计算 + all_to_all 通信的结果,数据类型与输入 x 保持一致,必须为 3 维。FLOAT16、BFLOAT16ND

需要特别留意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 取值,CheckTensorDimCommonShapeCheckTensorDimUniqueShape校验y_0 = x_0 * epx_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_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 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*EPC = 6*TPH = 2*TPM = 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 校验逻辑非常直观地映射了文档约束:

  • CheckTensorDimCommonShapey_0 == x_0 * epx_0 == w_0x_2 == w_1(不转置时 x 最后一维等于 weight 第 1 维);
  • CheckTensorDimUniqueShape:yShardType 为 0 时w_2 == y_2 * tpx_1 == y_1 * ep;为 1 时w_2 == y_2x_1 == y_1 * ep * tp
  • CheckShapeRangeM/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 == 1COverTp <= 640时启用 Lite 模式(isLite),此时输入 A 在 BMM 前需要先转置,kernel 会额外用 workspace 承载转置缓冲。

tiling 最终通过UpdateTilingKeyyShardFlag/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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/18 9:34:38

dlt 数据伪匿名化实战:使用 add_map 与加盐哈希隐藏 PII 列

dlt 数据伪匿名化实战&#xff1a;使用 add_map 与加盐哈希隐藏 PII 列 【免费下载链接】dlt data load tool (dlt) is an open source Python library that makes data loading easy &#x1f6e0;️ 项目地址: https://gitcode.com/GitHub_Trending/dl/dlt 伪匿名化&…

作者头像 李华
网站建设 2026/9/18 9:34:27

CS2自用命令指南:从autoexec到手感性能网络配置全解析

1. 从 CSGO 迁到 CS2&#xff1a;我为什么重新规整了一套自用命令CS2自用命令这个说法&#xff0c;听起来像是一份“自己用着顺手”的配置表&#xff0c;但它其实是我和游戏之间的一套固定对话方式。CS2正式上线后&#xff0c;我第一件事就是把用了很多年的 autoexec.cfg 从 CS…

作者头像 李华
网站建设 2026/9/18 9:32:55

DirectX 12 从设备初始化到贴图三角形:25 集调试路线

两年前我把网上流传最广的那套 DirectX 12 入门教程从头到尾抄了一遍&#xff1a;编译通过&#xff0c;窗口弹出来&#xff0c;三角形也确实画出来了。但从那之后整整三个星期&#xff0c;我卡在同一堆 Validation Error 里出不来——改分辨率黑屏、贴图加载出来一片漆黑、偶尔…

作者头像 李华
网站建设 2026/9/18 9:32:01

深度学习交通标志识别毕设:GTSRB、CNN分类与YOLO检测系统

1. 选题逻辑与整体架构&#xff1a;交通标志识别为什么是性价比最高的毕业设计方向带过几届学生的毕设之后&#xff0c;我对选题这件事有个很朴素的判断标准&#xff1a;数据能不能拿到、算法有没有公开基线、系统能不能跑起来给别人看。交通标志识别这个题目恰好三条全占。GTS…

作者头像 李华
网站建设 2026/9/18 9:31:04

Agent-Reach:高可靠Agent消息触达的幂等、重试与可观测设计

"Agent 能听懂话"和"Agent 的话能真的送到人手里"&#xff0c;中间隔着一整个工程世界。我前年参与过一个内部智能助手的落地&#xff0c;模型侧调得很顺&#xff0c;评测集上的准确率也好看&#xff0c;结果灰度上线第一周&#xff0c;用户投诉最多的问题…

作者头像 李华