资讯动态

CANN ops-transformer 中 AlltoAllvQuantGroupedMatMul 算子详解:路由专家 AlltoAllv 与量化 GroupedMatMul 的通信计算融合

发布时间:2026/9/18 14:29:35 来源:尧图企业网站定制
CANN ops-transformer 中 AlltoAllvQuantGroupedMatMul 算子详解路由专家 AlltoAllv 与量化 GroupedMatMul 的通信计算融合【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer导读AlltoAllvQuantGroupedMatMul 是 CANN ops-transformermc2/allto_allv_quant_grouped_mat_mul中面向 MoEMixture of Experts专家并行EP训练场景的高性能融合算子它将路由专家的 AlltoAllv 通信、permute 重排与量化 GroupedMatMul 计算融合为单个算子并与共享专家的量化 MatMul 并行执行整体遵循先通信后计算的执行模式。读完本文你将掌握该算子的计算流程、两段式 aclnn 接口的完整参数语义、pertensor/mx 两种量化模式的类型约束、groupSize 量化分组的编码规则与推导原理并能够结合仓库示例代码编写可运行的调用程序。算子定位MoE 专家并行中的通信-计算瓶颈在 MoE 大模型训练中路由expert parallelism阶段需要把每个 token 按照路由结果分发到不同卡上的专家再在专家维度执行分组矩阵乘GroupedMatMul。传统实现中AlltoAllv 集合通信与矩阵乘是两次独立的算子下发中间伴随 permute 重排存在明显的数据搬运与 kernel 启动开销。AlltoAllvQuantGroupedMatMul 将这条链路融合为一个算子从源码结构看该算子在 op_graph/allto_allv_quant_grouped_mat_mul_gen_task_training.cpp 与 op_graph/fallback_allto_allv_quant_grouped_mat_mul.cpp 中有对应的任务生成与回退路径路由专家路径gmmX先经过 AlltoAllv 通信与 permute得到本卡实际负责的 token 数据再按专家维度e 个专家做量化 GroupedMatMul共享专家路径mmX/mmWeight与本卡共享专家矩阵做量化 MatMul且与通信过程并行执行从而把通信等待时间隐藏在计算中。产品支持情况当前仓库中该算子仅支持Ascend 950DT其余产品系列均不支持产品是否支持Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品×Atlas A2 训练系列产品/Atlas A2 推理系列产品×Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×从 op_host/allto_allv_quant_grouped_mat_mul_def.cpp 可以看到算子仅注册了ascend950的 AICore 配置OpAICoreConfig aicore_config_950kernel 侧对应 op_kernel/arch35/allto_allv_quant_grouped_mat_mul_apt.cpp这也印证了硬件约束。功能与计算公式先通信后计算算子功能概括为完成路由专家 AlltoAllv、量化 GroupedMatMul 融合并实现与共享专家量化 MatMul 并行融合先通信后计算。假设通信域总卡数为epWorldSize每张卡上通信后路由专家个数为e每张卡的分组矩阵乘只负责本卡专家的计算则单卡上的完整计算过程如下本卡共享专家分组矩阵乘计算与通信并行$$ mm_y(mm_x \times mm_x_scale) (mm_weight \times mm_weight_scale) $$AlltoAllv 通信与 permute$$ permute_outAlltoAllv(gmm_x) $$本卡路由专家按专家维度分组矩阵乘计算$$ gmm_y(permute_out \times gmm_x_scale) (gmm_weight \times gmm_weight_scale) $$注意第 1 步与第 2、3 步是并行的共享专家的计算不依赖跨卡通信结果可以利用 AlltoAllv 通信的等待时间完成计算这是该算子性能设计的核心。两段式接口原型每个 aclnn 算子分为两段式接口必须先调用GetWorkspaceSize接口完成入参校验并计算所需 workspace 大小再调用执行接口完成计算。接口一V1 版本完整原型见 docs/aclnnAlltoAllvQuantGroupedMatMul.mdaclnnStatus aclnnAlltoAllvQuantGroupedMatMulGetWorkspaceSize( const aclTensor* gmmX, // 路由专家输入通信后作为分组乘左矩阵 const aclTensor* gmmWeight, // 路由专家分组乘右矩阵 const aclTensor* gmmXScale, // gmmX 量化系数 const aclTensor* gmmWeightScale, // gmmWeight 量化系数 const aclTensor* sendCountsTensorOptional,// 预留当前必须传 nullptr const aclTensor* recvCountsTensorOptional,// 预留当前必须传 nullptr const aclTensor* mmXOptional, // 可选共享专家左矩阵 const aclTensor* mmWeightOptional, // 可选共享专家右矩阵 const aclTensor* mmXScaleOptional, // 可选mmX 量化系数 const aclTensor* mmWeightScaleOptional, // 可选mmWeight 量化系数 int64_t gmmXQuantMode, // 量化模式1pertensor6mx int64_t gmmWeightQuantMode, int64_t mmXQuantMode, int64_t mmWeightQuantMode, const char* group, // 专家并行通信域名长度(0,128) int64_t epWorldSize, // ep 通信域大小 const aclIntArray* sendCounts, // 发送给其他卡的 token 数 const aclIntArray* recvCounts, // 接收其他卡的 token 数 bool transGmmWeight, // gmmWeight 是否转置 bool transMmWeight, // mmWeight 是否转置 int64_t groupSize, // 量化分组值编码 bool permuteOutFlag, // 是否输出 permute 结果 aclTensor* gmmY, // 输出路由专家计算结果 aclTensor* mmYOptional, // 输出共享专家计算结果 aclTensor* permuteOutOptional, // 输出permute 结果 uint64_t* workspaceSize, // 输出workspace 大小 aclOpExecutor** executor) // 输出op 执行器aclnnStatus aclnnAlltoAllvQuantGroupedMatMul( void* workspace, // Device 侧申请的 workspace 内存 uint64_t workspaceSize, // 由第一段接口获得 aclOpExecutor* executor, // 第一段接口返回的执行器 aclrtStream stream); // 任务执行 Stream接口二V2 版本见 docs/aclnnAlltoAllvQuantGroupedMatMulV2.mdV2 与 V1 相比仅新增一个commMode参数用于显式指定通信引擎支持ai_cpu与ccu两种取值其余参数与约束完全一致。参数详解与 shape 语义下表汇总了各输入/输出/属性参数的语义、数据类型与 shape 约束详细版可查阅两篇接口文档参数名输入/输出/属性描述数据类型数据格式gmmX输入进行 AlltoAllv 通信后结果作为 GroupedMatMul 左矩阵仅支持 2 维(BSK, H1)HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2、FLOAT4_E2M1NDgmmWeight输入GroupedMatMul 右矩阵仅支持 3 维不转置(e, H1, N1)转置(e, N1, H1)同 gmmXNDgmmXScale输入gmmX 量化系数pertensor 为 1 维(1)mx 为 3 维(BSK, ceildiv(H1,64), 2)FLOAT32、FLOAT8_E8M0NDgmmWeightScale输入gmmWeight 量化系数pertensor 为(1)mx 不转置(e, ceildiv(H1,64), N1, 2)转置(e, N1, ceildiv(H1,64), 2)FLOAT32、FLOAT8_E8M0NDsendCountsTensorOptional / recvCountsTensorOptional输入预留参数当前版本仅支持传 nullptr--mmXOptional输入共享专家左矩阵2 维(BS, H2)须与 mmWeightOptional 同时传入或同为 nullptr与 gmmX 一致NDmmWeightOptional输入共享专家右矩阵2 维不转置(H2, N2)转置(N2, H2)与 gmmWeight 一致NDmmXScaleOptional输入mmX 量化系数pertensor(1)mx(BS, ceildiv(H2,64), 2)FLOAT32、FLOAT8_E8M0NDmmWeightScaleOptional输入mmWeight 量化系数pertensor(1)mx 不转置(ceildiv(H2,64), N2, 2)转置(N2, ceildiv(H2,64), 2)FLOAT32、FLOAT8_E8M0NDgmmXQuantMode 等 4 个量化模式输入当前支持 1pertensor、6mxINT64-group输入专家并行通信域名字符串长度(0, 128)通过HcclGetCommName(HcclComm comm, char* commName)获取STRING-epWorldSize输入ep 通信域大小Ascend 950DT 支持 2/4/8/16/32/64/128/256INT64-sendCounts / recvCounts输入发送/接收各卡的 token 数INT64 元素长度e * epWorldSize最大 256需为 listaclIntArray*-transGmmWeight / transMmWeight输入右矩阵是否转置BOOL-groupSize输入量化分组值编码见下文INT64-permuteOutFlag输入permuteOutOptional 是否需要输出BOOL-gmmY输出路由专家计算结果2 维(A, N1)FLOAT16、BFLOAT16NDmmYOptional输出共享专家计算结果2 维(BS, N2)仅传入共享专家输入时输出与 gmmY 一致NDpermuteOutOptional输出permute 结果2 维(A, H1)仅 permuteOutFlag 为 true 时输出与 gmmX 一致ND量化模式枚举四个 QuantMode 参数共用0非量化、1pertensor、2perchannel、3pertoken、4pergroup、5perblock、6mx 量化、7pertoken 动态量化。当前版本接口仅开放1和6即 pertensor 量化与 mx 量化。返回值返回aclnnStatus状态码具体参见 aclnn 返回码。第一段接口完成入参校验典型错误包括返回值错误码场景ACLNN_ERR_PARAM_NULLPTR161001必选输入/输出或必选属性传入了空指针ACLNN_ERR_PARAM_INVALID161002gmmX、gmmWeight、mmXOptional 等的数据类型、数据格式或维度不在支持范围内量化模式与类型约束当前版本支持pertensor 量化与mx 量化两种模式其张量类型组合如下pertensor 量化QuantMode1gmmXgmmWeightgmmXScalegmmWeightScalemmXScalemmWeightScalegmmYHIFLOAT8HIFLOAT8FLOAT32FLOAT32FLOAT32FLOAT32FLOAT16/BFLOAT16mx 量化QuantMode6gmmXgmmWeightgmmXScalegmmWeightScalemmXScalemmWeightScalegmmYFLOAT8_E4M3FNFLOAT8_E4M3FNFLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16FLOAT8_E4M3FNFLOAT8_E5M2FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16FLOAT8_E5M2FLOAT8_E4M3FNFLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16FLOAT8_E5M2FLOAT8_E5M2FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16FLOAT4_E2M1FLOAT4_E2M1FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16此外类型一致性要求mmX与gmmX类型一致mmWeight与gmmWeight类型一致mmY与gmmY类型一致permuteOut与gmmX类型一致。这些约束在 op_host/allto_allv_quant_grouped_mat_mul_def.cpp 的算子定义中有完整登记gmm_x/gmm_weight支持 12 种量化组合对应的数据类型输出gmm_y/mm_y为 FLOAT16/BFLOAT16permute_out与gmm_x同类型。约束说明shape 变量定义BSK本卡发送的 token 数即 sendCounts 累加之和取值范围(0, 52428800)H1路由专家 hidden size(0, 65536)H2共享专家 hidden size(0, 12288]e单卡专家个数(0, 32]且e * epWorldSize最大 256N1路由专家的 head_num(0, 65536)N2共享专家的 head_num(0, 65536)BSbatch sequence sizeK选取 TopK 个专家范围[2, 8]A本卡收到的 token 数即 recvCounts 累加之和守恒关系ep 通信域内所有卡的 A 累加和等于所有卡的 BSK 累加和即 AlltoAllv 通信总量守恒。通信引擎约束V1 接口仅支持 AI_CPU 通信V2 接口支持 CCU 通信和 AI_CPU 通信。CCU 仅支持单机 UB 域内互联AI_CPU 可支持跨机 UB 域内互联。确定性计算aclnnAlltoAllvQuantGroupedMatMul与aclnnAlltoAllvQuantGroupedMatMulV2均为默认确定性实现。其他关键约束FLOAT4_E2M1 特殊约束mx 量化且 gmmX 与 gmmWeight 为 FLOAT4_E2M1 时H1 和 H2 必须为偶数且不能为 2同时当 transGmmWeight 和 transMmWeight 为 false 时N1 和 N2 必须为偶数。转置一致性gmmWeight与gmmWeightScale的转置状态必须一致同时转置或同时不转置mmWeight与mmWeightScale同理。op_api 层有专门的 CheckWeightScaleTransposeConsistency 校验逻辑。groupSize量化分组编码与推导规则groupSize用于表示量化中每个 scale 值在对应维度方向上可以覆盖多少个被量化数据即量化分组大小。它由三个方向的groupSizeM、groupSizeN、groupSizeK拼装而成每项占 16 位$$ groupSize groupSizeK ;|; groupSizeN \ll 16 ;|; groupSizeM \ll 32 $$使用要点仅当gmmXScale/gmmWeightScale/mmXScale/mmWeightScale输入都是 2 维及以上数据时 groupSize 才有效其他场景需传入0。传入的 groupSize 内部会先分解为 M/N/K 三个方向的值当其中有 1 个或多个为 0时接口会根据各输入 shape 重新推导对应方向的分组值groupSizeM 0groupSizeM m / scaleM需保证 m 能被 scaleM 整除其中 m 取自 gmmX/mmX 的 m 维scaleM 取自 gmmXScale/mmXScale 的 m 维groupSizeK 0groupSizeK k / scaleKk 与 gmmX/mmX 的 k 维一致scaleK 与 gmmXScale/mmXScale 的 k 维一致groupSizeN 0groupSizeN n / scaleNn 与 gmmWeight/mmWeight 的 n 维一致scaleN 与 gmmWeightScale/mmWeightScale 的 n 维一致。当满足重设条件且所有 scale 输入都是 2 维及以上、数据类型均为FLOAT8_E8M0时[groupSizeM, groupSizeN, groupSizeK]会统一推导为[1, 1, 32]对应 groupSize 值为4295032864。在示例代码中pertensor 量化、scale 为 1 维(1)groupSize直接传入 0即不启用分组量化。调用示例2 卡 aclnn 调用仓库提供了可直接参考的完整示例 examples/test_aclnn_allto_allv_quant_grouped_mat_mul.cpp其调用流程覆盖ACL 初始化 → 多卡 Context/Stream 创建 →HcclCommInitAll建立通信域 → 每卡一线程执行算子 → 同步等待 → 资源释放。核心 shape 配置如下示例以 2 卡 pertensor 量化为例constexpr int64_t EP_WORLD_SIZE 2; constexpr int64_t BS 4096; // batch sequence size constexpr int64_t K 2; // TopK 专家数 constexpr int64_t H 7168; // hidden size constexpr int64_t e 4; // 单卡专家个数 constexpr int64_t N1 4096; // 路由专家 head_num constexpr int64_t N2 4096; // 共享专家 head_num constexpr int64_t A BS * K; // 本卡接收 token 数 // 各张量 shape std::vectorint64_t gmmXShape {BS * K, H}; // 路由专家输入 (BSK, H1) std::vectorint64_t gmmWShape {e, H, N1}; // 路由专家权重 (e, H1, N1) std::vectorint64_t gmmYShape {A, N1}; // 路由专家输出 (A, N1) std::vectorint64_t permuteShape {A, H}; // permute 输出 (A, H1) std::vectorint64_t mmXShape {BS, H}; // 共享专家输入 (BS, H2) std::vectorint64_t mmWShape {H, N2}; // 共享专家权重 (H2, N2) std::vectorint64_t mmYShape {BS, N2}; // 共享专家输出 (BS, N2) std::vectorint64_t scaleShape {1}; // pertensor 缩放因子 // sendCounts/recvCounts每卡均分长度 e * epWorldSize 8 std::vectorint64_t sendCountsList(EP_WORLD_SIZE * e, BS * K / (EP_WORLD_SIZE * e)); std::vectorint64_t recvCountsList(EP_WORLD_SIZE * e, BS * K / (EP_WORLD_SIZE * e));调用算子时通过HcclGetCommName获取通信域名作为group参数量化模式全部取1pertensortransGmmWeight/transMmWeight为 falsegroupSize为 0permuteOutFlag为 trueret aclnnAlltoAllvQuantGroupedMatMulGetWorkspaceSize( gmmX, gmmW, gmmXScale, gmmWScale, nullptr, // sendCountsTensorOptional预留 nullptr, // recvCountsTensorOptional预留 mmX, mmW, mmXScale, mmWScale, 1, 1, 1, 1, // 四个量化模式均为 pertensor hcomName, EP_WORLD_SIZE, sendCounts, recvCounts, false, false, // transGmmWeight / transMmWeight groupSize, // pertensor 场景为 0 true, // permuteOutFlag gmmY, mmY, permute, workspaceSize, executor); if (workspaceSize 0) { aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } ret aclnnAlltoAllvQuantGroupedMatMul(workspaceAddr, workspaceSize, executor, args.stream); aclrtSynchronizeStreamWithTimeout(args.stream, 10000000);示例中的关键实现要点包括HIFLOAT8 数据用uint8_t模拟存储、缩放因子用 FLOAT32 填充 1.0、CreateAclTensor通过 shape 自动计算 strides 并以 ND 格式创建 tensor、main中按EP_WORLD_SIZE创建多线程分别驱动各 rank 执行。注意该量化接口仅支持 Ascend 950DT示例以 2 卡为准实际环境请按卡数修改EP_WORLD_SIZEHcclCommInitAll初始化的是设备连续编号的默认通信域多机场景需相应调整。源码级原理解析算子定义op_defallto_allv_quant_grouped_mat_mul_def.cpp 定义了 10 个输入4 个必选 6 个可选、3 个输出1 个必选 2 个可选以及group、ep_world_size、send_counts、recv_counts、trans_gmm_weight、comm_mode等属性。其中comm_mode属性默认值为ai_cpu并且注册了HcclGroup({group})说明该算子的集合通信是基于group指定的通信域发起的。入参校验与转置处理op_api在 op_api/aclnn_allto_allv_quant_grouped_mat_mul.cpp 中可以观察到该接口的实现细节必选参数校验gmmX、gmmWeight、gmmY、gmmXScale、gmmWeightScale不可为空量化模式必须是 1 或 6CheckNotNull预留参数约束sendCountsTensorOptional/recvCountsTensorOptional必须传 nullptrpermuteOutFlag与permuteOutOptional是否为 null 必须一致CheckNullStatus非连续 tensor 处理接口通过 stride 交换实现视图级转置如TransGmmWeightTensor支持转置的 gmmWeight 以非连续 tensor 传入当非连续与转置同时生效时判定为错误用法直接报错通信引擎选择V1 内部固定使用ai_cpu模式commMode ai_cpuV2 则把该字符串开放为入参支持ai_cpu/ccu。在 arch35DAV_3510上执行阶段会根据用户句柄中的 comm mode 调用NnopbaseSetHcclServerType设置 AI_CPU 或 CCU 通信服务类型aclnnAlltoAllvQuantGroupedMatMul 执行入口。shape 推导infershapeallto_allv_quant_grouped_mat_mul_infershape.cpp 实现了输出 shape 的推导逻辑与接口文档中的 shape 语义完全对应gmm_y2 维第一维 A 由 recvCounts 在e * epWorldSize长度上累加得到第二维为 N1转置时取gmmWeight第 1 维否则取第 2 维mm_y2 维(BS, N2)N2 依据 transMmWeight 从 mmWeight 的第 0/1 维选取permute_out2 维(A, H1)仅当 permuteOutFlag 为 true 时设置。这也说明 sendCounts/recvCounts 的累加一致性A 与 BSK 的守恒关系是 shape 推导正确的数据前提。测试与 golden 验证仓库在 tests/assets 提供了多设备 ACLNN 与 torch_npu E2E 的 TestSpec 适配spec.pygolden 参考实现位于 tests/assets/impl/golden.py支持expTokenNums每专家 token 数驱动的golden_gmm_alltoallv与默认golden_alltoallv_gmm两种正确性基准并支持 cascade三级流水结果的交叉校验。UT 侧还包含 op_api 的参数化测试含 nullptr 场景与 V2 场景以及 op_host 的 tiling 与 infershape 单测可用于验证不同量化模式、不同 shape 组合下的行为。小结AlltoAllvQuantGroupedMatMul 是 CANN ops-transformer 中面向 MoE 专家并行的高价值融合算子通过先通信后计算与共享专家 MatMul 并行化的设计将 AlltoAllv、permute 与量化 GroupedMatMul 合并为一次下发。理解其两段式接口的参数语义尤其 groupSize 的编码与推导、量化模式与类型约束、转置一致性要求是正确使用的前提。更完整的约束与逐参数说明建议直接查阅 README.md、aclnnAlltoAllvQuantGroupedMatMul 接口文档 与 V2 接口文档并结合 调用示例 进行实践。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价