FEATURED · 精选文章

CANN ops-transformer GroupedMatMulAlltoAllv 算子实战:路由专家计算与 AlltoAllv 通信的融合方案

发布时间 / 2026/9/20 12:01:19
来源 / 创域科博编辑部
栏目 / 资讯中心
CANN ops-transformer GroupedMatMulAlltoAllv 算子实战:路由专家计算与 AlltoAllv 通信的融合方案 算子库人工智能深度学习Ascend【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载导读GroupedMatMulAlltoAllv是 CANN ops-transformer 项目中面向 MoEMixture 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_93Atlas A3 系列与ascend910bAtlas A2 系列共用aicore_config_a3ascend950Ascend 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、BFLOAT16NDgmmWeight输入GroupedMatMul 计算的右矩阵数据类型与 gmmX 保持一致支持 3 维shape 为 (e, H1, N1)FLOAT16、BFLOAT16NDsendCountsTensorOptional输入可选输入shape 为 (e × epWorldSize,)当前版本暂不支持传 nullptrINT32、INT64NDrecvCountsTensorOptional输入可选输入shape 为 (e × epWorldSize,)当前版本暂不支持传 nullptrINT32、INT64NDmmXOptional输入可选输入共享专家 MatMul 的左矩阵需与 mmWeightOptional 同时传入或同为 nullptr数据类型与 gmmX 保持一致支持 2 维shape 为 (BS, H2)FLOAT16、BFLOAT16NDmmWeightOptional输入可选输入共享专家 MatMul 的右矩阵需与 mmXOptional 同时传入或同为 nullptr数据类型与 gmmX 保持一致支持 2 维shape 为 (H2, N2)FLOAT16、BFLOAT16NDgroup输入专家并行的通信域名称字符串长度要求 (0, 128)STRINGNDepWorldSize输入EP 通信域 sizeAtlas A2 系列支持 2、4、8Atlas A3 系列支持 8、16、32、64、128Ascend 950DT 支持 2、4、8、16、32、64INT64NDsendCounts输入表示发送给其他卡的 token 数元素类型 INT64取值大小为 e × epWorldSizeAIV 通信最大为 1024其他通信引擎最大为 256aclIntArray*元素类型 INT64NDrecvCounts输入表示接收其他卡的 token 数元素类型 INT64取值大小为 e × epWorldSizeAIV 通信最大为 1024其他通信引擎最大为 256aclIntArray*元素类型 INT64NDtransGmmWeight输入gmmWeight 是否需要转置true 表示需要转置false 表示不转置BOOLNDtransMmWeight输入共享专家 mmWeightOptional 是否需要转置true 表示需要转置false 表示不转置BOOLNDy输出最终计算结果数据类型与 gmmX 保持一致支持 2 维shape 为 (BSK, N1)FLOAT16、BFLOAT16NDmmYOptional输出共享专家 MatMul 的输出数据类型与 mmXOptional 保持一致支持 2 维shape 为 (BS, N2)仅当传入 mmXOptional 与 mmWeightOptional 时才输出FLOAT16、BFLOAT16ND参数语义补充说明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分派Arch35A5 架构且commMode ccu时走ccu_stream/CCU GenTask其余情况走kfc_streamAICPU 通信服务器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 最大支持 256N1路由专家的 head_num取值范围 (0, 65536)N2共享专家的 head_num取值范围 (0, 65536)BSbatch sequence sizeK选取 TopK 个专家。Atlas A3 系列产品的 AIV 通信支持 [2, 16]其他场景支持 [2, 8]A本卡发送的 token 数是 sendCounts 参数累加之和守恒关系EP 通信域内所有卡的 A 参数累加和等于所有卡上的 BSK 参数累加和。在 shape 推导实现中InferGMMOutputShape会以e × epWorldSize为长度校验sendCounts/recvCounts的 attr 数组大小并对recvCounts逐元素累加得到输出第一维BSKInferMMOutputShape则在三个可选输入均存在时推导共享专家输出(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_NULLPTR161001必选输入/输出/属性传了空指针、ACLNN_ERR_PARAM_INVALID161002gmmX、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申请 workspaceworkspaceSize 0时用aclrtMalloc在 Device 侧申请调用第二段接口执行计算同步并回收资源aclrtSynchronizeStreamWithTimeout等待任务结束随后销毁 Tensor、释放 Device 内存、销毁 stream/context/通信域并aclrtResetDevice、aclFinalize。核心调用示例多卡 EP 场景以下代码节选自仓库示例完整版见 examples/test_aclnn_grouped_mat_mul_allto_allv.cpp展示了单线程内完成一次融合计算的核心逻辑。示例配置为EP_WORLD_SIZE8、BS4096、K2、H7168、e4、N1N24096A 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::vectorint64_t gmmXShape {A, H}; std::vectorint64_t gmmWShape {e, H, N1}; std::vectorint64_t gmmYShape {BS * K, N1}; std::vectorint64_t mmXShape {BS, H}; std::vectorint64_t mmWShape {H, N2}; std::vectorint64_t mmYShape {BS, N2}; std::vectorint64_t sendCountsList(EP_WORLD_SIZE * e, A / (EP_WORLD_SIZE * e)); std::vectorint64_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 / commaclrtResetDevice ... return 0; }主函数中需要为 EP 域内的每张卡分别创建 context 与 stream通过HcclCommInitAll(EP_WORLD_SIZE, devices, comms)一次性初始化整个通信域然后为每个 rank 启动一个线程执行上述逻辑最后 join 所有线程并aclFinalize()。示例中用std::vectoruint16_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和aivAtlas A2 系列产品仅支持aiv不支持ai_cpu和ccuAscend 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 同 dtypeaclnn 封装与参数校验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 与 Kernelop_host/op_tiling含 arch22 的 MTE tiling 与 A3 tiling、arch35 的 A5 tiling与 op_kernel含 arch22 的 MTE kernel 与通用 A3 kernel支撑动态 shape 下的编译与多 kernel 执行测试用例tests/ut 覆盖 op_apiV1/V2 两段式接口调用、op_hostinfershape 与 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上加速计算。项目地址https://gitcode.com/cann/ops-transformer点击查看免费下载相关推荐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上一篇Objective-C-RSA 项目常见问题解决方案下一篇5分钟极速打包PakePlus让网页变身专业桌面应用创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻