CANN ops-transformer 中 AlltoAllvQuantGroupedMatMul 算子详解:路由专家 AlltoAllv 与量化 GroupedMatMul 的通信计算融合
2026/9/18 14:29:20 网站建设 项目流程

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

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

导读

AlltoAllvQuantGroupedMatMul 是 CANN ops-transformer(mc2/allto_allv_quant_grouped_mat_mul)中面向 MoE(Mixture 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_950),kernel 侧对应 op_kernel/arch35/allto_allv_quant_grouped_mat_mul_apt.cpp,这也印证了硬件约束。

功能与计算公式:先通信后计算

算子功能概括为:完成路由专家 AlltoAllv、量化 GroupedMatMul 融合,并实现与共享专家量化 MatMul 并行融合,先通信后计算

假设通信域总卡数为epWorldSize,每张卡上通信后路由专家个数为e,每张卡的分组矩阵乘只负责本卡专家的计算,则单卡上的完整计算过程如下:

  1. 本卡共享专家分组矩阵乘计算(与通信并行):

    $$ mm_y=(mm_x \times mm_x_scale) @ (mm_weight \times mm_weight_scale) $$

  2. AlltoAllv 通信与 permute

    $$ permute_out=AlltoAllv(gmm_x) $$

  3. 本卡路由专家按专家维度分组矩阵乘计算

    $$ gmm_y=(permute_out \times gmm_x_scale) @ (gmm_weight \times gmm_weight_scale) $$

注意第 1 步与第 2、3 步是并行的:共享专家的计算不依赖跨卡通信结果,可以利用 AlltoAllv 通信的等待时间完成计算,这是该算子性能设计的核心。

两段式接口原型

每个 aclnn 算子分为两段式接口:必须先调用GetWorkspaceSize接口完成入参校验并计算所需 workspace 大小,再调用执行接口完成计算。

接口一(V1 版本),完整原型见 docs/aclnnAlltoAllvQuantGroupedMatMul.md:

aclnnStatus 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, // 量化模式,1=pertensor,6=mx 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.md:V2 与 V1 相比仅新增一个commMode参数,用于显式指定通信引擎,支持ai_cpuccu两种取值,其余参数与约束完全一致。

参数详解与 shape 语义

下表汇总了各输入/输出/属性参数的语义、数据类型与 shape 约束(详细版可查阅两篇接口文档):

参数名输入/输出/属性描述数据类型数据格式
gmmX输入进行 AlltoAllv 通信后结果作为 GroupedMatMul 左矩阵,仅支持 2 维(BSK, H1)HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2、FLOAT4_E2M1ND
gmmWeight输入GroupedMatMul 右矩阵,仅支持 3 维:不转置(e, H1, N1),转置(e, N1, H1)同 gmmXND
gmmXScale输入gmmX 量化系数:pertensor 为 1 维(1);mx 为 3 维(BSK, ceildiv(H1,64), 2)FLOAT32、FLOAT8_E8M0ND
gmmWeightScale输入gmmWeight 量化系数:pertensor 为(1);mx 不转置(e, ceildiv(H1,64), N1, 2),转置(e, N1, ceildiv(H1,64), 2)FLOAT32、FLOAT8_E8M0ND
sendCountsTensorOptional / recvCountsTensorOptional输入预留参数,当前版本仅支持传 nullptr--
mmXOptional输入共享专家左矩阵,2 维(BS, H2),须与 mmWeightOptional 同时传入或同为 nullptr与 gmmX 一致ND
mmWeightOptional输入共享专家右矩阵,2 维:不转置(H2, N2),转置(N2, H2)与 gmmWeight 一致ND
mmXScaleOptional输入mmX 量化系数:pertensor(1);mx(BS, ceildiv(H2,64), 2)FLOAT32、FLOAT8_E8M0ND
mmWeightScaleOptional输入mmWeight 量化系数:pertensor(1);mx 不转置(ceildiv(H2,64), N2, 2),转置(N2, ceildiv(H2,64), 2)FLOAT32、FLOAT8_E8M0ND
gmmXQuantMode 等 4 个量化模式输入当前支持 1(pertensor)、6(mx)INT64-
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、BFLOAT16ND
mmYOptional输出共享专家计算结果,2 维(BS, N2),仅传入共享专家输入时输出与 gmmY 一致ND
permuteOutOptional输出permute 结果,2 维(A, H1),仅 permuteOutFlag 为 true 时输出与 gmmX 一致ND

量化模式枚举(四个 QuantMode 参数共用):0非量化、1pertensor、2perchannel、3pertoken、4pergroup、5perblock、6mx 量化、7pertoken 动态量化。当前版本接口仅开放16,即 pertensor 量化与 mx 量化。

返回值:返回aclnnStatus状态码,具体参见 aclnn 返回码。第一段接口完成入参校验,典型错误包括:

返回值错误码场景
ACLNN_ERR_PARAM_NULLPTR161001必选输入/输出或必选属性传入了空指针
ACLNN_ERR_PARAM_INVALID161002gmmX、gmmWeight、mmXOptional 等的数据类型、数据格式或维度不在支持范围内

量化模式与类型约束

当前版本支持pertensor 量化mx 量化两种模式,其张量类型组合如下:

pertensor 量化(QuantMode=1)

gmmXgmmWeightgmmXScalegmmWeightScalemmXScalemmWeightScalegmmY
HIFLOAT8HIFLOAT8FLOAT32FLOAT32FLOAT32FLOAT32FLOAT16/BFLOAT16

mx 量化(QuantMode=6)

gmmXgmmWeightgmmXScalegmmWeightScalemmXScalemmWeightScalegmmY
FLOAT8_E4M3FNFLOAT8_E4M3FNFLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16
FLOAT8_E4M3FNFLOAT8_E5M2FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16
FLOAT8_E5M2FLOAT8_E4M3FNFLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16
FLOAT8_E5M2FLOAT8_E5M2FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16
FLOAT4_E2M1FLOAT4_E2M1FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT8_E8M0FLOAT16/BFLOAT16

此外,类型一致性要求:mmXgmmX类型一致,mmWeightgmmWeight类型一致,mmYgmmY类型一致,permuteOutgmmX类型一致。这些约束在 op_host/allto_allv_quant_grouped_mat_mul_def.cpp 的算子定义中有完整登记:gmm_x/gmm_weight支持 12 种量化组合对应的数据类型,输出gmm_y/mm_y为 FLOAT16/BFLOAT16,permute_outgmm_x同类型。

约束说明

shape 变量定义

  • BSK:本卡发送的 token 数,即 sendCounts 累加之和,取值范围(0, 52428800)
  • H1:路由专家 hidden size,(0, 65536)
  • H2:共享专家 hidden size,(0, 12288]
  • e:单卡专家个数,(0, 32],且e * epWorldSize最大 256;
  • N1:路由专家的 head_num,(0, 65536)
  • N2:共享专家的 head_num,(0, 65536)
  • BS:batch sequence size;
  • K:选取 TopK 个专家,范围[2, 8]
  • A:本卡收到的 token 数,即 recvCounts 累加之和;
  • 守恒关系:ep 通信域内所有卡的 A 累加和等于所有卡的 BSK 累加和(即 AlltoAllv 通信总量守恒)。

通信引擎约束

  • V1 接口:仅支持 AI_CPU 通信;
  • V2 接口:支持 CCU 通信和 AI_CPU 通信。CCU 仅支持单机 UB 域内互联,AI_CPU 可支持跨机 UB 域内互联。

确定性计算

aclnnAlltoAllvQuantGroupedMatMulaclnnAlltoAllvQuantGroupedMatMulV2均为默认确定性实现

其他关键约束

  • FLOAT4_E2M1 特殊约束:mx 量化且 gmmX 与 gmmWeight 为 FLOAT4_E2M1 时,H1 和 H2 必须为偶数且不能为 2;同时当 transGmmWeight 和 transMmWeight 为 false 时,N1 和 N2 必须为偶数。
  • 转置一致性gmmWeightgmmWeightScale的转置状态必须一致(同时转置或同时不转置);mmWeightmmWeightScale同理。op_api 层有专门的 CheckWeightScaleTransposeConsistency 校验逻辑。

groupSize:量化分组编码与推导规则

groupSize用于表示量化中每个 scale 值在对应维度方向上可以覆盖多少个被量化数据(即量化分组大小)。它由三个方向的groupSizeMgroupSizeNgroupSizeK拼装而成,每项占 16 位:

$$ groupSize = groupSizeK ;|; groupSizeN \ll 16 ;|; groupSizeM \ll 32 $$

使用要点:

  1. 仅当gmmXScale/gmmWeightScale/mmXScale/mmWeightScale输入都是 2 维及以上数据时 groupSize 才有效,其他场景需传入0
  2. 传入的 groupSize 内部会先分解为 M/N/K 三个方向的值,当其中有 1 个或多个为 0时,接口会根据各输入 shape 重新推导对应方向的分组值:
    • groupSizeM = 0groupSizeM = m / scaleM(需保证 m 能被 scaleM 整除),其中 m 取自 gmmX/mmX 的 m 维,scaleM 取自 gmmXScale/mmXScale 的 m 维;
    • groupSizeK = 0groupSizeK = k / scaleK,k 与 gmmX/mmX 的 k 维一致,scaleK 与 gmmXScale/mmXScale 的 k 维一致;
    • groupSizeN = 0groupSizeN = n / scaleN,n 与 gmmWeight/mmWeight 的 n 维一致,scaleN 与 gmmWeightScale/mmWeightScale 的 n 维一致。
  3. 当满足重设条件且所有 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::vector<int64_t> gmmXShape = {BS * K, H}; // 路由专家输入 (BSK, H1) std::vector<int64_t> gmmWShape = {e, H, N1}; // 路由专家权重 (e, H1, N1) std::vector<int64_t> gmmYShape = {A, N1}; // 路由专家输出 (A, N1) std::vector<int64_t> permuteShape = {A, H}; // permute 输出 (A, H1) std::vector<int64_t> mmXShape = {BS, H}; // 共享专家输入 (BS, H2) std::vector<int64_t> mmWShape = {H, N2}; // 共享专家权重 (H2, N2) std::vector<int64_t> mmYShape = {BS, N2}; // 共享专家输出 (BS, N2) std::vector<int64_t> scaleShape = {1}; // pertensor 缩放因子 // sendCounts/recvCounts:每卡均分,长度 e * epWorldSize = 8 std::vector<int64_t> sendCountsList(EP_WORLD_SIZE * e, BS * K / (EP_WORLD_SIZE * e)); std::vector<int64_t> recvCountsList(EP_WORLD_SIZE * e, BS * K / (EP_WORLD_SIZE * e));

调用算子时通过HcclGetCommName获取通信域名作为group参数,量化模式全部取1(pertensor),transGmmWeight/transMmWeight为 false,groupSize为 0,permuteOutFlag为 true:

ret = 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_def)

allto_allv_quant_grouped_mat_mul_def.cpp 定义了 10 个输入(4 个必选 + 6 个可选)、3 个输出(1 个必选 + 2 个可选)以及groupep_world_sizesend_countsrecv_countstrans_gmm_weightcomm_mode等属性。其中comm_mode属性默认值为"ai_cpu",并且注册了HcclGroup({"group"}),说明该算子的集合通信是基于group指定的通信域发起的。

入参校验与转置处理(op_api)

在 op_api/aclnn_allto_allv_quant_grouped_mat_mul.cpp 中可以观察到该接口的实现细节:

  • 必选参数校验gmmXgmmWeightgmmYgmmXScalegmmWeightScale不可为空,量化模式必须是 1 或 6(CheckNotNull);
  • 预留参数约束sendCountsTensorOptional/recvCountsTensorOptional必须传 nullptr,permuteOutFlagpermuteOutOptional是否为 null 必须一致(CheckNullStatus);
  • 非连续 tensor 处理:接口通过 stride 交换实现"视图级转置"(如TransGmmWeightTensor),支持转置的 gmmWeight 以非连续 tensor 传入;当"非连续与转置同时生效"时判定为错误用法直接报错;
  • 通信引擎选择:V1 内部固定使用"ai_cpu"模式(commMode = "ai_cpu");V2 则把该字符串开放为入参,支持ai_cpu/ccu。在 arch35(DAV_3510)上,执行阶段会根据用户句柄中的 comm mode 调用NnopbaseSetHcclServerType设置 AI_CPU 或 CCU 通信服务类型(aclnnAlltoAllvQuantGroupedMatMul 执行入口)。

shape 推导(infershape)

allto_allv_quant_grouped_mat_mul_infershape.cpp 实现了输出 shape 的推导逻辑,与接口文档中的 shape 语义完全对应:

  • gmm_y:2 维,第一维 A 由 recvCounts 在e * epWorldSize长度上累加得到,第二维为 N1(转置时取gmmWeight第 1 维,否则取第 2 维);
  • mm_y:2 维(BS, N2),N2 依据 transMmWeight 从 mmWeight 的第 0/1 维选取;
  • permute_out:2 维(A, H1),仅当 permuteOutFlag 为 true 时设置。

这也说明 sendCounts/recvCounts 的累加一致性(A 与 BSK 的守恒关系)是 shape 推导正确的数据前提。

测试与 golden 验证

仓库在 tests/assets 提供了多设备 ACLNN 与 torch_npu E2E 的 TestSpec 适配(spec.py),golden 参考实现位于 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询