- 算子库
- 人工智能
- 深度学习
- Ascend
【免费下载链接】ops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
导读
GroupedMatMulAlltoAllv是 CANN ops-transformer 项目中面向 MoE(Mixture of Experts)大模型场景的融合算子:它将路由专家的 GroupedMatMul、Unpermute 与 AlltoAllv 集合通信融合为单个算子,同时把共享专家的 MatMul 计算并行叠加进来,整体遵循"先计算后通信"的执行策略。阅读本文后,你将掌握该算子的产品支持范围、输入输出参数语义、shape 与通信约束,以及基于 aclnn 两段式接口编写多卡(EP 专家并行)调用程序并完成编译运行的完整方法。
产品支持情况
根据 mc2/grouped_mat_mul_allto_allv/README.md 中的支持矩阵,该算子适用于以下产品:
| 产品 | 是否支持 |
|---|---|
| Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
从算子注册文件 grouped_mat_mul_allto_allv_def.cpp 可以看到,该算子为不同的昇腾架构配置了独立的 AICore 计算配置:ascend910_93(Atlas A3 系列)与ascend910b(Atlas A2 系列)共用aicore_config_a3,ascend950(Ascend 950DT)使用aicore_config_a5,并且在算子级通过.MC2().HcclGroup({"group"})将group属性声明为 HCCL 集合通信组,这也是该算子属于 MC2(融合通信计算)算子族的核心标志。
功能与计算原理
融合了什么
该算子完成三件事的融合:
- 路由专家 GroupedMatMul:按专家维度分组执行矩阵乘;
- Unpermute:将 GroupedMatMul 的输出按路由结果重排回原始 token 顺序;
- AlltoAllv:在专家并行(EP)通信域内执行变长 all-to-all 通信,把属于其他卡的数据发送过去。
与此同时,共享专家的MatMul计算被设计为与上述链路并行执行,通过计算与通信的重叠隐藏延迟。官方定义为"先计算后通信":通信的数据来自本卡本地计算完成后的结果,从而避免先通信再计算的串行等待。
计算公式
路由专家链路:
$$ gmmY = gmmX \times gmmWeight \ unpermuteOut = Unpermute(gmmY) \ y = AlltoAllv(unpermuteOut) $$
共享专家链路:
$$ mmY = mmX \times mmWeight $$
两路计算在同一个算子内完成,最终同时产出y(路由专家最终输出)与可选的mmYOptional(共享专家输出)。
数据流视角
gmmX是输入 token 经过 router 挑选后的激活,其第一维A表示本卡需要发送给各 EP 卡的 token 总数;gmmWeight以(e, H1, N1)的三维结构承载单卡上的e个路由专家权重;- 通信阶段,每张卡通过
sendCounts/recvCounts描述与通信域内各卡的收发 token 数量; - 最终每卡收到的 token 总数记为
BSK,输出y的 shape 为(BSK, N1)。
从接口头文件 aclnn_grouped_mat_mul_allto_allv.h 的注释可以进一步印证数据关系:e表示单卡上的专家数量,A = recvCounts的累加和(注:按输出 shape 推导的实际含义,A为 sendCounts 累加和、BSK为 recvCounts 累加和,详见下文约束说明),且 EP 通信域内所有卡的A累加和等于所有卡的BSK累加和,这正是 AlltoAllv 通信守恒关系。
参数说明
算子的全部参数如下表(摘自 README.md 参数说明,并补充了 aclnn 接口的维度/连续 Tensor 约束):
| 参数名 | 输入/输出 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| gmmX | 输入 | 该输入进行 AlltoAllv 通信,通信后结果作为 GroupedMatMul 计算的左矩阵,支持 2 维,shape 为 (A, H1) | FLOAT16、BFLOAT16 | ND |
| gmmWeight | 输入 | GroupedMatMul 计算的右矩阵,数据类型与 gmmX 保持一致,支持 3 维,shape 为 (e, H1, N1) | FLOAT16、BFLOAT16 | ND |
| sendCountsTensorOptional | 输入 | 可选输入,shape 为 (e × epWorldSize,),当前版本暂不支持,传 nullptr | INT32、INT64 | ND |
| recvCountsTensorOptional | 输入 | 可选输入,shape 为 (e × epWorldSize,),当前版本暂不支持,传 nullptr | INT32、INT64 | ND |
| mmXOptional | 输入 | 可选输入,共享专家 MatMul 的左矩阵,需与 mmWeightOptional 同时传入或同为 nullptr,数据类型与 gmmX 保持一致,支持 2 维,shape 为 (BS, H2) | FLOAT16、BFLOAT16 | ND |
| mmWeightOptional | 输入 | 可选输入,共享专家 MatMul 的右矩阵,需与 mmXOptional 同时传入或同为 nullptr,数据类型与 gmmX 保持一致,支持 2 维,shape 为 (H2, N2) | FLOAT16、BFLOAT16 | ND |
| group | 输入 | 专家并行的通信域名称,字符串长度要求 (0, 128) | STRING | ND |
| epWorldSize | 输入 | EP 通信域 size:Atlas A2 系列支持 2、4、8;Atlas A3 系列支持 8、16、32、64、128;Ascend 950DT 支持 2、4、8、16、32、64 | INT64 | ND |
| sendCounts | 输入 | 表示发送给其他卡的 token 数,元素类型 INT64,取值大小为 e × epWorldSize,AIV 通信最大为 1024,其他通信引擎最大为 256 | aclIntArray*(元素类型 INT64) | ND |
| recvCounts | 输入 | 表示接收其他卡的 token 数,元素类型 INT64,取值大小为 e × epWorldSize,AIV 通信最大为 1024,其他通信引擎最大为 256 | aclIntArray*(元素类型 INT64) | ND |
| transGmmWeight | 输入 | gmmWeight 是否需要转置,true 表示需要转置,false 表示不转置 | BOOL | ND |
| transMmWeight | 输入 | 共享专家 mmWeightOptional 是否需要转置,true 表示需要转置,false 表示不转置 | BOOL | ND |
| y | 输出 | 最终计算结果,数据类型与 gmmX 保持一致,支持 2 维,shape 为 (BSK, N1) | FLOAT16、BFLOAT16 | ND |
| mmYOptional | 输出 | 共享专家 MatMul 的输出,数据类型与 mmXOptional 保持一致,支持 2 维,shape 为 (BS, N2),仅当传入 mmXOptional 与 mmWeightOptional 时才输出 | FLOAT16、BFLOAT16 | ND |
参数语义补充说明
- gmmX 与 gmmWeight 的维度校验:在 shape 推导实现 grouped_mat_mul_allto_allv_infershape.cpp 中,
CheckDims强制要求 gmmX 为 2 维、gmmWeight 为 3 维,并校验 MatMul 内维匹配(不转置时要求 gmmWeight 的H1维等于 gmmX 的H1维;转置时则取 gmmWeight 最后一维),不满足时直接报 "Dim of gmmX and dim of gmmWeight do not match for MatMul"。 - 共享专家三件套必须同传同缺:在 aclnn_grouped_mat_mul_allto_allv.cpp 的
CheckNullStatus中,mmXOptional、mmWeightOptional、mmYOptional要么全为 nullptr,要么全非空,混传会返回ACLNN_ERR_PARAM_INVALID并记录 "should all be null or all not be null" 的日志。 - sendCounts/recvCounts 不允许为空:aclnn 第一段接口会通过
CheckSendAndRecv校验aclIntArray非空且元素个数大于 0。 - 测试侧的输入约定:仓库测试资产 tests/assets/inputs.py 也复现了同样的规则——
send_counts与recv_counts长度必须一致、ep_world_size必须为正、mm_x/mm_weight必须成对出现,可作为编写调用时的输入校验参考。
约束说明
通信引擎约束
不同产品支持不同的集合通信引擎(即 AlltoAllv 由哪个引擎执行):
- Atlas A2 训练/推理系列:仅支持 AIV 通信;
- Atlas A3 训练/推理系列:支持 AI_CPU 通信和 AIV 通信;
- Ascend 950DT:支持 CCU 通信和 AI_CPU 通信。其中 CCU 仅支持单机 UB 域内互联,AI_CPU 可支持跨机 UB 域内互联。
这一约束在算子定义与图侧 GenTask 中有直接对应:comm_mode属性默认值为"ai_cpu"(见 grouped_mat_mul_allto_allv_def.cpp),而在 grouped_mat_mul_allto_allv_gen_task_training.cpp 中,任务生成会根据目标架构与comm_mode分派:Arch35(A5 架构)且commMode == "ccu"时走ccu_stream/CCU GenTask,其余情况走kfc_stream(AICPU 通信服务器)GenTask。aclnn 封装层 aclnn_grouped_mat_mul_allto_allv.cpp 同样根据平台架构调用NnopbaseSetHcclServerType设置 AICPU 或 CCU 通信服务器类型。
shape 变量的取值范围
- BSK:本卡接收的 token 数,是 recvCounts 参数累加之和,取值范围 (0, 52428800);
- H1:路由专家 hidden size 隐藏层大小,取值范围 (0, 65536);
- H2:共享专家 hidden size 隐藏层大小,取值范围 (0, 12288];
- e:单卡上专家个数。AIV 通信要求 e > 0 且 e × epWorldSize 最大支持 1024;其他通信引擎要求 e ≤ 32 且 e × epWorldSize 最大支持 256;
- N1:路由专家的 head_num,取值范围 (0, 65536);
- N2:共享专家的 head_num,取值范围 (0, 65536);
- BS:batch sequence size;
- K:选取 TopK 个专家。Atlas A3 系列产品的 AIV 通信支持 [2, 16],其他场景支持 [2, 8];
- A:本卡发送的 token 数,是 sendCounts 参数累加之和;
- 守恒关系:EP 通信域内所有卡的 A 参数累加和等于所有卡上的 BSK 参数累加和。
在 shape 推导实现中,InferGMMOutputShape会以e × epWorldSize为长度校验sendCounts/recvCounts的 attr 数组大小,并对recvCounts逐元素累加得到输出第一维BSK;InferMMOutputShape则在三个可选输入均存在时推导共享专家输出(BS, N2)。若 shape 未知(动态 shape),相关维度会先置为 -1,由运行时二次推导。
Atlas A2 的 HCCL_BUFFSIZE 配置
Atlas A2 训练/推理系列产品上,A和BSK均需在 [1, 5000000] 范围内,N1不超过 32768。此外,通信域内各卡的HCCL_BUFFSIZE需按最大发送量设置,满足:
HCCL_BUFFSIZE >= max(200, ceil(A * N1 * 2 / 1048576) + 21)单位是 MiB。其中 FLOAT16 和 BFLOAT16 每个元素均占 2 字节,21 MiB 为控制区预留空间。V2 接口文档 aclnnGroupedMatMulAlltoAllvV2.md 还进一步说明:sendCounts/recvCounts是 INT64 直接计数数组,按[rank][localExpert]顺序展平,长度均须等于e × epWorldSize,元素为非负数,分别不超过 A 和 BSK,且累加和分别等于 A 和 BSK。
性能提示
Atlas A3 训练/推理系列产品上,单卡通信量在 2MB 以下可能存在性能劣化,规划模型规模时应注意控制通信数据量。
调用方式:aclnn 两段式接口
该算子的官方调用方式是 aclnn 接口,遵循"两段式"调用范式:先调用aclnnGroupedMatMulAlltoAllvGetWorkspaceSize完成入参校验并计算 workspace 大小,再调用aclnnGroupedMatMulAlltoAllv在指定 stream 上执行计算。完整样例位于 examples/test_aclnn_grouped_mat_mul_allto_allv.cpp 和 docs/aclnnGroupedMatMulAlltoAllv.md。
函数原型
aclnnStatus aclnnGroupedMatMulAlltoAllvGetWorkspaceSize( const aclTensor* gmmX, const aclTensor* gmmWeight, const aclTensor* sendCountsTensorOptional, const aclTensor* recvCountsTensorOptional, const aclTensor* mmXOptional, const aclTensor* mmWeightOptional, const char* group, int64_t epWorldSize, const aclIntArray* sendCounts, const aclIntArray* recvCounts, bool transGmmWeight, bool transMmWeight, aclTensor* y, aclTensor* mmYOptional, uint64_t* workspaceSize, aclOpExecutor** executor)aclnnStatus aclnnGroupedMatMulAlltoAllv( void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream)第一段接口还会返回错误码:ACLNN_ERR_PARAM_NULLPTR(161001,必选输入/输出/属性传了空指针)、ACLNN_ERR_PARAM_INVALID(161002,gmmX、gmmWeight、group、epWorldSize、sendCounts、recvCounts 等参数的数据类型、数据格式或维度不在支持范围内)。
调用流程
标准调用步骤为:
- 初始化环境:
aclInit→ 每张卡aclrtSetDevice→aclrtCreateContext→aclrtCreateStream; - 初始化集合通信域:通过 HCCL 的
HcclCommInitAll创建 EP 通信域,并用HcclGetCommName取出通信域名字符串作为group入参; - 构造 Tensor 与 counts:用
aclCreateTensor创建 ND 格式 Tensor,用aclCreateIntArray构造sendCounts/recvCounts(长度EP_WORLD_SIZE * e,示例中每个元素取A / (EP_WORLD_SIZE * e)的平均分配值); - 调用第一段接口获取
workspaceSize与executor; - 申请 workspace:
workspaceSize > 0时用aclrtMalloc在 Device 侧申请; - 调用第二段接口执行计算;
- 同步并回收资源:
aclrtSynchronizeStreamWithTimeout等待任务结束,随后销毁 Tensor、释放 Device 内存、销毁 stream/context/通信域并aclrtResetDevice、aclFinalize。
核心调用示例(多卡 EP 场景)
以下代码节选自仓库示例(完整版见 examples/test_aclnn_grouped_mat_mul_allto_allv.cpp),展示了单线程内完成一次融合计算的核心逻辑。示例配置为EP_WORLD_SIZE=8、BS=4096、K=2、H=7168、e=4、N1=N2=4096,A = BS * K = 8192:
#include "acl/acl.h" #include "hccl/hccl.h" #include "aclnnop/aclnn_grouped_mat_mul_allto_allv.h" // shape 基本信息 constexpr int64_t EP_WORLD_SIZE = 8; constexpr int64_t BS = 4096; constexpr int64_t K = 2; constexpr int64_t H = 7168; constexpr int64_t e = 4; constexpr int64_t N1 = 4096; constexpr int64_t N2 = 4096; constexpr int64_t A = BS * K; // 本卡发送 token 数 int LaunchOneThreadAlltoAllvGmm(Args &args) { int ret = aclrtSetCurrentContext(args.context); char hcomName[128] = {0}; ret = HcclGetCommName(args.hcclComm, hcomName); // 取得通信域名作为 group std::vector<int64_t> gmmXShape = {A, H}; std::vector<int64_t> gmmWShape = {e, H, N1}; std::vector<int64_t> gmmYShape = {BS * K, N1}; std::vector<int64_t> mmXShape = {BS, H}; std::vector<int64_t> mmWShape = {H, N2}; std::vector<int64_t> mmYShape = {BS, N2}; std::vector<int64_t> sendCountsList(EP_WORLD_SIZE * e, A / (EP_WORLD_SIZE * e)); std::vector<int64_t> recvCountsList(EP_WORLD_SIZE * e, A / (EP_WORLD_SIZE * e)); // ... 通过 CreateAclTensor 构造 gmmX/gmmW/gmmY/mmX/mmW/mmY 六个 aclTensor ... aclIntArray *sendCounts = aclCreateIntArray(sendCountsList.data(), sendCountsList.size()); aclIntArray *recvCounts = aclCreateIntArray(recvCountsList.data(), recvCountsList.size()); // 调用第一阶段接口:校验入参并计算 workspace 大小 ret = aclnnGroupedMatMulAlltoAllvGetWorkspaceSize(gmmX, gmmW, nullptr, nullptr, mmX, mmW, hcomName, EP_WORLD_SIZE, sendCounts, recvCounts, false, false, gmmY, mmY, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, return ret); if (workspaceSize > 0) { ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, return ret); } // 调用第二阶段接口:执行计算 ret = aclnnGroupedMatMulAlltoAllv(workspaceAddr, workspaceSize, executor, args.stream); CHECK_RET(ret == ACL_SUCCESS, return ret); // 同步等待任务执行结束 ret = aclrtSynchronizeStreamWithTimeout(args.stream, 10000000); // ... 释放 Tensor / Device 内存 / stream / context / comm,aclrtResetDevice ... return 0; }主函数中需要为 EP 域内的每张卡分别创建 context 与 stream,通过HcclCommInitAll(EP_WORLD_SIZE, devices, comms)一次性初始化整个通信域,然后为每个 rank 启动一个线程执行上述逻辑,最后 join 所有线程并aclFinalize()。示例中用std::vector<uint16_t>以 2 字节元素承载 FLOAT16/BFLOAT16 数据(gmmXHostData、gmmWHostData等),并用aclCreateTensor按 ND 格式、行主序 strides 创建张量。
提示:本示例还调用了部分 HCCL 集合通信库接口(
HcclGetCommName、HcclCommInitAll、HcclCommDestroy),具体编译与运行样例的方法可参考仓库 docs 中的样例编译运行说明。
V2 接口:显式指定通信引擎
针对不同产品通信引擎能力差异,仓库同时提供 V2 接口aclnnGroupedMatMulAlltoAllvV2,其文档见 docs/aclnnGroupedMatMulAlltoAllvV2.md。与 V1 接口相比,核心变更是新增commMode参数,让用户显式指定当前使用的通信引擎:
- Atlas A3 系列产品:支持
ai_cpu和aiv; - Atlas A2 系列产品:仅支持
aiv,不支持ai_cpu和ccu; - Ascend 950DT:支持
ai_cpu和ccu。
V2 的函数原型在 V1 的基础上于group之后插入const char* commMode,其余参数与两段式调用流程保持一致。此外,V2 支持的产品范围更广:README 中 V1 的产品矩阵将 Atlas A2 系列列为支持,而 V2 文档明确列出的支持项为 Ascend 950DT、Atlas A3 与 Atlas A2 训练/推理系列,epWorldSize在 Atlas A2 上支持 2、4、8。V2 示例(docs/aclnnGroupedMatMulAlltoAllvV2.md 中的调用示例)在 V1 基础上仅需在GetWorkspaceSize调用中多传一个"ai_cpu"字符串:
ret = aclnnGroupedMatMulAlltoAllvV2GetWorkspaceSize(gmmX, gmmW, sendCountsTensor, recvCountsTensor, mmX, mmW, hcomName, "ai_cpu", EP_WORLD_SIZE, sendCounts, recvCounts, false, false, gmmY, mmY, &workspaceSize, &executor);源码级实现路径
若希望深入理解该算子的落地方式,可按如下路径阅读仓库源码:
- 算子注册与属性定义:op_host/grouped_mat_mul_allto_allv_def.cpp —— 定义 6 个输入(gmm_x、gmm_weight、send_counts_tensor、recv_counts_tensor、mm_x、mm_weight)、2 个输出(y、mm_y)以及 group、ep_world_size、send_counts、recv_counts、trans_gmm_weight、trans_mm_weight、comm_mode 等属性;
- shape 推导:op_host/grouped_mat_mul_allto_allv_infershape.cpp —— 校验维度与 MatMul 内维匹配、累加 recvCounts 推导 BSK、推导 mmY shape,并完成输出数据类型透传(输出与 gmmX 同 dtype);
- aclnn 封装与参数校验:op_api/aclnn_grouped_mat_mul_allto_allv.cpp 与 op_api/aclnn_grouped_mat_mul_allto_allv.h —— 两段式接口实现、空指针与 counts 校验、按架构设置 HCCL 通信服务器类型;
- 图侧任务生成:op_graph/grouped_mat_mul_allto_allv_gen_task_training.cpp —— 根据架构与 comm_mode 选择 CCU 或 AICPU 的 GenTask 与 stream 类型;
- Tiling 与 Kernel:op_host/op_tiling(含 arch22 的 MTE tiling 与 A3 tiling、arch35 的 A5 tiling)与 op_kernel(含 arch22 的 MTE kernel 与通用 A3 kernel),支撑动态 shape 下的编译与多 kernel 执行;
- 测试用例:tests/ut 覆盖 op_api(V1/V2 两段式接口调用)、op_host(infershape 与 tiling)与 op_kernel 的单元测试,tests/assets/inputs.py 提供输入参数的合法性校验与 shape 调整逻辑,可作为实现对齐的参考基准。
其中算子定义中的jitCompile.flag = "static_false"与multiKernelSupportDynamicGraph.value = "multi_kernel"等扩展配置说明该算子以动态 shape、多 kernel 的方式运行,二进制可被复用。
结语
GroupedMatMulAlltoAllv 是 CANN ops-transformer 中针对 MoE 专家并行训练的典型 MC2 融合算子:它以"先计算后通信"的方式把路由专家的 GroupedMatMul、Unpermute、AlltoAllv 与共享专家 MatMul 合并执行,兼顾了计算与通信的重叠以及算子级图融合带来的调度收益。开发者在使用时需要重点核对三类信息:一是目标产品的支持矩阵与可用通信引擎(AIV/AI_CPU/CCU),二是epWorldSize、e、A、BSK、K等参数之间的取值范围与守恒关系,三是 Atlas A2 上HCCL_BUFFSIZE的配置要求。按照本文给出的两段式接口流程与仓库示例,即可快速完成多卡环境下的集成与验证。
- 算子库
- 人工智能
- 深度学习
- Ascend
【免费下载链接】ops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
相关推荐
CANN ops-transformer 中 AlltoAllvQuantGroupedMatMul 算子详解:路由专家 AlltoAllv 与量化 GroupedMatMul 的通信计算融合
CANN ops transformer 中 AlltoAllvQuantGroupedMatMul 算子详解:路由专家 AlltoAllv 与量化 Group
算子库人工智能深度学习AscendCANN ops-transformer 中 aclnnAlltoAllvQuantGroupedMatMulV2 算子:AlltoAllv 通信与量化 GroupedMatMul 融合实战指南
CANN ops transformer 中 aclnnAlltoAllvQuantGroupedMatMulV2 算子:AlltoAllv 通信与量化 Gro
算子库人工智能深度学习AscendCANN ops-transformer AlltoAllvGroupedMatMul 算子深度解析:路由专家通信与计算融合的 aclnn 两段式接口实战
CANN ops transformer AlltoAllvGroupedMatMul 算子深度解析:路由专家通信与计算融合的 aclnn 两段式接口实战 本文
算子库人工智能深度学习Ascend
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考