尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

解读 ATB GroupedMatmulWithRouting 算子:MoE 场景下带路由的分组矩阵乘实现

解读 ATB GroupedMatmulWithRouting 算子:MoE 场景下带路由的分组矩阵乘实现 解读 ATB GroupedMatmulWithRouting 算子MoE 场景下带路由的分组矩阵乘实现【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost导读GroupedMatmulWithRouting 是 CANN ascend-transformer-boostATB推理算子库中面向 MoEMixture-of-Experts稀疏专家路由场景的高性能分组矩阵乘算子用于将每个 token 按路由结果选出的 topK 个专家权重执行矩阵乘法实现专家网络的 Up/Down 投影。本文以仓库知识条目 grouped_matmul_with_routing/index.md 及其 路由文件 为骨架结合 Operation、Runner、Kernel 三层源码与测试用例完整梳理该算子的参数语义、Tensor 规格、形状约束、实现链路与验证方式帮助你快速读懂并正确使用这一算子。1. 算子定位MoE 稀疏路由的分组矩阵乘GroupedMatmulWithRouting 属于 ATB 推理infer类算子复杂度评级为 M其核心能力在参数头文件中定义如下实现了 GroupedMatmulWithRouting 算子的 Up 和 Down 方法将 topK 个专家权重与 token 激活值做矩阵乘法计算。infer_op_params.h在 MoE 推理中Gating 网络为每个 token 选出 topK 个专家随后 token 激活值需要与这些被选中专家的权重分别做矩阵乘。若逐专家循环执行普通 Matmul会产生大量小算子启动开销与中间搬运该算子将按路由索引取专家权重 分组矩阵乘融合为一个整体一次调用完成所有被选中专家的计算并按路由顺序重排输出。从源码结构看该算子由三层协作完成Operation 层GroupedMatmulWithRoutingOperation负责参数校验、形状推导InferShape与 Runner 创建Runner 层GroupedMatmulWithRoutingRunnerOpsRunner 子类负责把输入 Tensor 组装成底层 MoeGmm 内核图KernelGraphKernel 层MoeGmmOperation及其 kernel 实现位于src/kernels/mixkernels/moe_gmm/完成实际的分组矩阵乘与可选的反量化计算。知识路由文件将该算子标记为Runner 类型: OpsRunner, Operation、ACLNN: no即它不走 ACLNN 适配路径而是直接以 Operation OpsRunner 的组合注册执行这与上述三层结构一致。注路由文件中标注的 Kernel 目录为src/kernels/mixkernels/laser_attention而依据当前仓库源码实际检索MoeGmm 内核实现位于 src/kernels/mixkernels/moe_gmm/阅读时请以实际源码路径为准。1.1 适用硬件前提该算子不是通用算子使用前必须满足平台约束。在CreateOperation入口处有显式的平台检查if (!GetSingletonConfig().Is910B()) { ATB_LOG(ERROR) only support Atlas 800I A2 inference product; return ERROR_INVALID_PARAM; }对应 grouped_matmul_with_routing_operation.cpp参数头文件也以\warning明确仅 Atlas 800I A2 推理产品支持该算子。在高层测试用例中SocVersion 均为Ascend910B而 Ascend310P 平台用例预期返回ERROR_INVALID_PARAM进一步印证了这一限制。2. 参数结构体与字段语义算子参数定义于 infer_op_params.h 的GroupedMatmulWithRoutingParam共四个业务字段加一段预留空间字段类型默认值语义与取值约束groupedMatmulTypeGroupedMatmulTypeint 枚举GROUPED_MATMUL_UP执行 Up0还是 Down1投影transposeBbooltrue是否转置 B 矩阵专家权重true表示权重以[numExperts, hiddenOut, hiddenIn]排布topKint32_t0必须显式给出每个 token 选中的专家个数合法范围 [2, 10]outDataTypeaclDataTypeACL_DT_UNDEFINED输出数据类型ACL_DT_UNDEFINED表示非量化量化场景取ACL_FLOAT16/ACL_BF16rsv[16]uint8_t全 0预留参数其中枚举定义enum GroupedMatmulType : int { GROUPED_MATMUL_UP 0, //! 默认值。up类型。 GROUPED_MATMUL_DOWN //! down类型。 };2.1 字段的底层影响topK 决定中间形状的缩放Up 阶段每个 token 会产出 topK 个中间结果输出第一维为token 数 × topKDown 阶段把中间结果合并回 token 数输出第一维为token 数 ÷ topK详见第 5 节 InferShape。transposeB 决定权重的隐藏维排布transposeB true时专家权重 shape 为[numExperts, hiddenOut, hiddenIn]即 K 维在最后天然适配以 token 为 M 维的矩阵乘为false时权重为[numExperts, hiddenIn, hiddenOut]。Runner 组装内核图时也依据该字段从权重张量的第 1 或第 2 维解析hiddenSize[1]。outDataType 决定是否进入量化链路ACL_DT_UNDEFINED走 4 输入的非量化路径权重 FP16/BF16输出与输入同 dtypeACL_FLOAT16/ACL_BF16走 6 输入的 W8A8 量化反量化路径权重 INT8输出为指定精度的反量化结果。3. 输入输出 Tensor 规格算子的输入输出规格定义在 ops_configs/atb_ops_info.ini分为普通与量化两套算子配置。3.1 非量化GroupedMatmulWithRoutingOperation槽位名称dtypeformat说明input0inputfloat16 / bf16ndtoken 激活值shape[tokens, hiddenIn]input1weightfloat16 / bf16nd专家权重shape[numExperts, hiddenOut, hiddenIn]transposeBtrue或[numExperts, hiddenIn, hiddenOut]input2ecountsint32nd每个专家分到的 token 数expert countsshape[numExperts]input3indicesint32nd每个槽位对应的专家索引Up 时长度tokens × topKDown 时长度tokensoutput0outputfloat16 / bf16nd计算结果3.2 量化W8A8GroupedMatmulWithRoutingQuantOperation槽位名称dtypeformat说明input0inputint8nd量化后的 token 激活值input1weightint8nd / fractal_nz量化专家权重NZ 格式时会触发专门的视图重排见 6.2 节input2ecountsint32nd专家 token 计数input3indicesint32nd专家路由索引input4scalefloatnd权重侧反量化 scaleshape[numExperts, hiddenOut]input5scale2floatnd激活侧反量化 scaleshape[tokens]output0outputfloat16 / bf16nd反量化输出对应到 Operation 内部输入槽位以常量索引管理grouped_matmul_with_routing_operation.cppstatic const uint32_t IN_TENSOR_ACTENSOR 0; // token 激活值 static const uint32_t IN_TENSOR_EXPERTWEIGHT 1; // 专家权重 static const uint32_t IN_TENSOR_EXPERTCOUNT 2; // 专家 token 计数 static const uint32_t IN_TENSOR_EXPERTINDEX 3; // 专家索引 static const uint32_t IN_TENSOR_NSCALE 4; // 权重侧 scale量化 static const uint32_t IN_TENSOR_MSCALE 5; // 激活侧 scale量化输入数量由outDataType决定非量化IN_TENSOR_NUM 4量化QUANT_IN_TENSOR_NUM 6输出恒为OUT_TENSOR_NUM 1GetInputNum/GetOutputNum。4. 源码结构与推荐阅读顺序算子源码位于 src/ops/ops_infer/grouped_matmul_with_routing/共 4 个文件与路由文件 grouped_matmul_with_routing.md 的文件清单一致#文件角色重点关注1grouped_matmul_with_routing_operation.hOperation 定义输入输出数量、InferShape 签名、参数成员2grouped_matmul_with_routing_operation.cppOperation 实现CreateOperation()参数/平台校验、InferShapeImpl()、CreateRunner()决策3grouped_matmul_with_routing_runner.hRunner 头文件OpsRunner 继承关系、SetupKernelGraph()签名4grouped_matmul_with_routing_runner.cppRunner 实现内核图组装、反量化类型映射、NZ 视图重排配套阅读路径参数定义include/atb/infer_op_params.h底层内核src/kernels/mixkernels/moe_gmm/含moe_gmm_operation.cpp、moe_gmm_kernel.cpp、op_kernel/moe_gmm.cce、op_kernel/moe_gmm_w8a8.cce、tiling 目录算子配置ops_configs/atb_ops_info.ini测试用例tests/high_level_test/GroupedMatmulWithRoutingOperation/ 与 tests/apitest/opstest/python/operations/groupedmatmulwithrouting/推荐按 头文件 → Operation 实现 → Runner 实现 → Kernel 的顺序阅读与路由文件给出的阅读顺序一致。5. Operation 层校验与形状推导5.1 创建入口的校验链CreateOperationgrouped_matmul_with_routing_operation.cpp依次执行三类检查任一失败即返回ERROR_INVALID_PARAM空指针检查operation输出指针为空平台检查仅支持 Atlas 800I A2Is910B()参数范围检查topK必须在 [2, 10] 之间TOPK_MIN_VALUE 2、TOPK_MAX_VALUE 10outDataType仅允许ACL_DT_UNDEFINED、ACL_FLOAT16、ACL_BF16三者之一。5.2 InferShapeUp 放大、Down 还原InferShapeImplgrouped_matmul_with_routing_operation.cpp的核心逻辑if (param_.groupedMatmulType GROUPED_MATMUL_UP) { outTensorDescs.at(0) inTensorDescs.at(0); outTensorDescs.at(0).shape.dims[0] inTensorDescs.at(0).shape.dims[0] * param_.topK; } else { outTensorDescs.at(0) inTensorDescs.at(0); outTensorDescs.at(0).shape.dims[0] inTensorDescs.at(0).shape.dims[0] / param_.topK; } outTensorDescs.at(0).shape.dims[1] OperationUtil::GetYTensorN(inTensorDescs.at(1), param_.transposeB); if (param_.outDataType ! ACL_DT_UNDEFINED) { outTensorDescs.at(0).dtype param_.outDataType; }要点第一维Up 阶段输出行数为tokens × topK每个 token 对应 topK 个中间结果Down 阶段输出行数为tokens ÷ topK合并回 token 粒度。第二维取专家权重的输出隐藏维hiddenOut其取值受transposeB影响——transposeB true时取权重第 1 维否则取第 2 维与OutTensorDimCheck中的计算一致。量化场景输出 dtype 被覆盖为outDataTypeFP16/BF16这是INT8 计算 反量化输出的体现。5.3 SetupCheck运行时张量形状校验SetupCheckImpl与InferShapeCheckImpl会在运行前对每个输入做形状级校验违规返回ERROR_INVALID_TENSOR_DIM。各校验函数与规则如下校验函数校验对象规则TokenExpertTensorCheckinput0 与 weightinput0 第 2 维hiddenIn必须等于权重隐藏输入维权重 hiddenIn/hiddenOut 必须32 对齐Up 时 hiddenIn ∈ [32, 5120]、hiddenOut ∈ [32, 256]Down 时 hiddenIn ∈ [32, 256]、hiddenOut ∈ [32, 5120]ExpertCountTensorCheckecounts 与 weightecounts 第 0 维专家数必须等于权重第 0 维且专家数 ∈ [128, 256]ExpertIndexTensorCheckinput0 与 indicesUp 时 indices 长度 tokens × topKDown 时 indices 长度 tokens激活 token 数 ∈ [128, 512]NScaleTensorCheckscale 与 weightscale 第 0 维 专家数第 1 维 权重 hiddenOutMScaleTensorCheckscale2 与 input0scale2 第 0 维 tokens这些范围常量定义在 grouped_matmul_with_routing_operation.cpp如WEIGHT_MIN_VALUE 128、WEIGHT_MAX_VALUE 256、ACTIVATION_MIN_VALUE 128、ACTIVATION_MAX_VALUE 512、ALIGHMENT_NUMBER 32等。这意味着该算子面向的是中大规模 MoE 推理场景如 128~256 个专家、128~512 个激活 token小规模测试需注意下限约束。5.4 Runner 创建决策CreateRunnergrouped_matmul_with_routing_operation.cpp恒返回GroupedMatmulWithRoutingRunner不依赖运行上下文做额外分支但构造 Operation 时已根据outDataType选择不同的 IR 配置GroupedMatmulWithRoutingQuantOperation或GroupedMatmulWithRoutingOperation对应 ini 中的两套算子配置。6. Runner 层组装 MoeGmm 内核图GroupedMatmulWithRoutingRunner继承自OpsRunner核心方法是SetupKernelGraphgrouped_matmul_with_routing_runner.cpp它在一次调用中完成组图将 4 或 6 个输入组织成一个单节点内核图6.1 内核图组装与参数映射kernelGraph_.nodes.resize(1); auto moeGmmNode kernelGraph_.nodes.at(0); AtbOps::OpParam::MoeGmm opParam; opParam.moeGmmDequantType AtbOps::OpParam::MoeGmm::NO_DEQUANT; if (param_.outDataType ACL_BF16) { opParam.moeGmmDequantType AtbOps::OpParam::MoeGmm::DEQ_BF16; } else if (param_.outDataType ACL_FLOAT16) { opParam.moeGmmDequantType AtbOps::OpParam::MoeGmm::DEQ_FP16; } opParam.moeGmmMode static_castAtbOps::OpParam::MoeGmm::MoeGmmMode(param_.groupedMatmulType); opParam.transposeB param_.transposeB ? 1 : 0; opParam.topK static_castuint32_t(param_.topK); opParam.hiddenSize.at(0) inputTensor.desc.dims.at(1); opParam.hiddenSize.at(1) param_.transposeB ? weightTensor.desc.dims.at(1) : weightTensor.desc.dims.at(2); moeGmmNode.opDesc {0, MoeGmmOperation, opParam};映射关系一目了然moeGmmDequantType由outDataType决定ACL_BF16 → DEQ_BF16、ACL_FLOAT16 → DEQ_FP16、否则NO_DEQUANTmoeGmmMode与groupedMatmulType枚举值直接对应0 为 UP、1 为 DOWNhiddenSize[0] 激活值隐藏维hiddenInhiddenSize[1] 权重隐藏输出维hiddenOut受 transposeB 影响取第 1 或第 2 维。6.2 量化路径与 NZ 视图重排量化场景outDataType ! ACL_DT_UNDEFINED时输入从 4 个扩展为 6 个多出的weightscale与activatescale对应 ini 中的 scale / scale2Mki::Tensor weightscale kernelGraph_.inTensors.at(inTensorNum); Mki::Tensor activatescale kernelGraph_.inTensors.at(inTensorNum); moeGmmNode.inTensors {inputTensor, weightTensor, ecountTensor, indiceIntensor, weightscale, activatescale};特别地若量化权重为FRACTAL_NZ格式isWeightNz会对权重张量注册一个视图变换函数把 3 维 NZ 权重重排为 4 维int64_t align INT8_ALIGN; // 32 if (oldDims.size() SIZE_3) { newDims {oldDims.at(0), UtilsInternal::AlignUp(oldDims.at(DIM_2), align) / align, UtilsInternal::AlignUp(oldDims.at(1), DEFAULT_ALIGN), align}; }即以[专家数, K 维按 32 对齐分块, N 维按 16 对齐, 32]的布局描述 NZ 权重的实际内存排布供底层 kernel 直接消费。最后通过REG_RUNNER_TYPE(GroupedMatmulWithRoutingRunner)与REG_OP_PARAM(AtbOps::OpParam::MoeGmm)完成注册串起 Operation → Runner → Kernel 的调用链。7. Kernel 层MoeGmmOperation 的计算语义底层内核由 src/kernels/mixkernels/moe_gmm/moe_gmm_operation.cpp 定义MoeGmmOperation输入数量随反量化类型切换NO_DEQUANT为 4DEQ_FP16/DEQ_BF16为 6GetInputNum与上层 Runner 的组图结果一致输出数量恒为 1。其InferShapeImplmoe_gmm_operation.cpp与上层 Operation 的推导保持同一语义case OpParam::MoeGmm::MOE_GMM_UP: outDims.emplace_back(index.dims[0]); // tokens × topK outDims.emplace_back(attrs.hiddenSize[1]); // hiddenOut break; case OpParam::MoeGmm::MOE_GMM_DOWN: tensorDescOut inTensorDescA; tensorDescOut.dims[0] tensorDescOut.dims[0] / attrs.topK; tensorDescOut.dims[1] attrs.hiddenSize[1]; break;反量化场景下输出 dtype 被覆盖为TENSOR_DTYPE_BF16或TENSOR_DTYPE_FLOAT16。实际计算内核位于 moe_gmm_kernel.cpp、op_kernel/moe_gmm.cceFP16/BF16 路径与 op_kernel/moe_gmm_w8a8.cceINT8 W8A8 路径tiling 计算见 tiling/moe_gmm_tiling.cpp。从文件组织可以推断非量化与量化各有一套独立 kernel 实现tiling 层负责根据 shape 分块。8. 形状约束速查表实操参考综合 Operation 层校验与 ini 配置使用该算子前建议对照以下约束维度/参数约束topK[2, 10]专家数weight 第 0 维 ecounts 第 0 维[128, 256]激活 token 数[128, 512]权重 hiddenIn / hiddenOut必须 32 对齐UphiddenIn[32, 5120]UphiddenOut[32, 256]DownhiddenIn[32, 256]DownhiddenOut[32, 5120]输入输出 dtype非量化 fp16/bf16量化 int8 输入 fp16/bf16 输出权重格式nd非量化nd / fractal_nz量化硬件平台仅 Atlas 800I A2Ascend910B推理产品一个典型合法的 Up 场景来自高层测试用例为input[128, 5120]、weight[160, 192, 5120]、ecounts[160]、indices[768]128 token × topK6topK6、transposeBtrue输出[768, 192]。9. 测试与验证用例覆盖分析仓库为该算子提供了两层测试覆盖正常路径与大量边界/错误路径9.1 高层测试 CSVtests/high_level_test/GroupedMatmulWithRoutingOperation/Dtype_dataFormat/GroupedMatmulWithRoutingOperation_TestCase.csv 以表格驱动方式定义了 60 个用例覆盖维度包括Up / Down 两种模式groupedMatmulType取 0/1 的正常用例如noquant_up_fp16_noerror、down_fp16_noerrordtype 组合fp16、bf16 非量化int8 scale/scale2 的 W8A8 量化upquant_fp16_noerror、upquant_bf16_noerrortransposeB 开关true与false两套权重排布格式组合nd / fractal_nz 的权重与输出量化 NZ 场景up_nz_quant_noerror错误注入topK越界1、11、未给出→ERROR_INVALID_PARAMhidden 维不匹配/不对齐 →ERROR_INVALID_TENSOR_DIM输入 dtype、format 与配置不符 →ERROR_INVALID_TENSOR_INI_MATCH非 910B 平台 →ERROR_INVALID_PARAM输出 dtype 与参数不一致 → Setup 阶段报错前缀S:等等。其中ExpectedError列的前缀可用于判断报错阶段C:为创建参数错误、I:为 InferShape/输入校验错误、S:为 Setup 阶段错误这为排查问题提供了明确线索。9.2 APITest Python 用例tests/apitest/opstest/python/operations/groupedmatmulwithrouting/ 下按方向与是否量化拆分为四个测试文件test_groupedmatmulwithrouting_up.py/test_groupedmatmulwithrouting_down.py非量化 Up / Downtest_groupedmatmulwithrouting_up_quant.py/test_groupedmatmulwithrouting_down_quant.pyW8A8 量化 Up / Down。对应 CSV 驱动用例位于 tests/apitest/opstest/csv/grouped_matmul_with_routing.csv测试整体运行方式可参考仓库的 测试框架指南 与 apitest/kernelstest/op_test.py。此外 tests/unittest/ops/test_op_param_size.cpp 中对该参数结构体的大小有单测覆盖确保参数结构在 ABI 层面稳定。10. 使用注意事项小结平台先行该算子仅支持 Atlas 800I A2Ascend910B推理产品其他平台调用会直接失败无需继续排查参数。topK 必须显式设置默认值为 0 不在合法区间未设置或越界2 或 10都会在创建阶段报ERROR_INVALID_PARAM。量化与非量化输入个数不同非量化 4 输入、量化 6 输入多出 scale、scale2且量化权重可选 fractal_nz 格式配合 Runner 的 NZ 视图重排使用。形状有下限与对齐要求专家数 128~256、激活 token 128~512、权重隐藏维 32 对齐——小 batch 或小模型的调试用例需要先满足这些下限约束否则会命中ERROR_INVALID_TENSOR_DIM。Up/Down 输出形状不对称Up 输出行数为tokens × topKDown 输出行数为tokens ÷ topK二者第二维均为权重隐藏输出维搭建 MoE 前向图时需按此衔接。【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表