资讯动态

ascend-transformer-boost GroupTopkOperation 路由与源码导读:MoE 分组 TopK 算子的实现与调试验证

发布时间:2026/9/18 7:01:50 来源:尧图企业网站定制
ascend-transformer-boost GroupTopkOperation 路由与源码导读MoE 分组 TopK 算子的实现与调试验证【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost导读本文围绕 CANN ascend-transformer-boost 仓库中的GroupTopkOperationMoE 路由场景下的分组 TopK算子展开基于 .agent/knowledge/routing/group_topk.md 路由文件梳理其代码组织、推荐阅读顺序与源码路径并结合 include/atb/infer_op_params.h 中的参数定义、src/ops/ops_infer/group_topk/ 的 Operation/Runner 实现以及 src/kernels/kernels/group_topk 的 Kernel 与 Tiling 代码讲解算子从参数校验、InferShape、KernelGraph 构建到 Kernel 派发、Tiling 切分的完整调用链并给出测试用例作为验证依据。读完本文你将掌握该算子的文件布局、每个参数groupNum/k/groupMultiFlag/n的取值约束与语义以及分组取 TopK、其余组置零这一 MoE 剪枝路由逻辑在源码中的落地方式。1. 算子定位分组 TopK 在 MoE 路由中的作用GroupTopkOperation属于推理infer类别算子复杂度等级为 S简单对应文件数量为 4 个运行时不走 ACLNN而是走OpsRunnerOperation路径。从 include/atb/infer_op_params.h 中的参数注释可以明确其计算语义将输入 inTensor0 中维度 1inTensor0 有 2 个维度维度 0 和维度 1的数据分groupNum个组每组取最大值然后选出每组最大值中前k个最后将非前k个组的数据全部置零。这一语义对应 MoEMixture of Experts推理中的专家剪枝/路由场景每个 token 在全部专家上的得分按组聚合后只保留 TopK 个组的专家其余专家的路由得分被置零从而为后续GroupedMatmulWithRouting等算子的稀疏计算做准备。测试代码 tests/apitest/kernelstest/group_topk/test_group_topk.py 中的典型用例512, 256, 8, 3, 2即 token512、expert256、groupNum8、k3、kInner2标注为 typical case from DeepseekV3印证了该算子面向 DeepSeek-V3 这类 MoE 大模型路由的实际来源。2. 文件清单与角色路由文件给出了该算子涉及的 4 个文件均位于 src/ops/ops_infer/group_topk/角色如下#文件角色1group_topk_operation.cppOperation 定义参数校验、InferShape、CreateRunner 决策2group_topk_operation.hOperation 定义类声明、输入输出数量、虚函数签名3group_topk_ops_runner.cppOps RunnerKernelGraph 构建、参数下传4group_topk_ops_runner.hOps Runner原生 Ops 执行接口声明值得说明的是这 4 个文件只是 ATB 侧算子封装层的文件。算子真正的计算逻辑还依赖 Kernel 侧目录 src/kernels/kernels/group_topk其中包含group_topk_kernel.cppKernel 能力检查、group_topk_operation.cppKernel 侧算子注册与 Kernel 选择、tiling/group_topk_tiling.cpp与tiling/group_topk_tiling.hTiling 计算、op_kernel/group_topk.cppAscendC 计算实现。阅读时可以按ATB 封装层 → Kernel 层的顺序递进。3. 推荐阅读顺序路由文件给出的推荐阅读顺序为顺序文件重点关注1group_topk_operation.h了解输入输出数量、InferShape 签名2group_topk_operation.cppCreateRunner() 决策逻辑3group_topk_ops_runner.h原生 Ops 执行接口4group_topk_ops_runner.cpp原生 Ops 调用链 平台适配结合源码按此顺序阅读的实际收益是group_topk_operation.h可以看到GroupTopkOperation继承自OperationBase构造函数接收infer::GroupTopkParam重写了GetInputNum()返回 2、GetOutputNum()返回 1、InferShapeImpl、InferShapeCheckImpl、SetupCheckImpl、CreateRunner、GetParamJson等接口。这些签名就是 ATB 算子框架对 Operation 的全部约定。group_topk_operation.cpp重点看模板特化CreateOperationGroupTopkParam中的逐项参数校验见下文第 6 节以及CreateRunner中直接std::make_sharedGroupTopkOpsRunner(param_)的决策逻辑。group_topk_ops_runner.hGroupTopkOpsRunner继承OpsRunner仅重写SetupKernelGraph和SetParam并定义了operator只比较groupNum与k用于运行时参数变更检测。group_topk_ops_runner.cppSetupKernelGraph中构造AsdOps::OpParam::GroupTopk原生参数并挂载到 KernelGraphNode 上完成 ATB 参数到 ASD 原生参数的映射即调用链 平台适配的落点。4. 源码路径总览Op 目录src/ops/ops_infer/group_topk/Kernel 目录src/kernels/kernels/group_topk参数头文件include/atb/infer_op_params.hGroupTopkParam定义于第 25242580 行5. 参数定义GroupTopkParam 详解GroupTopkParam在 include/atb/infer_op_params.h#L2532-L2580 中定义共 4 个有效字段加 1 个预留字段字段类型默认值取值范围说明groupNumint32_t1[1, 专家总数]每个 token 的分组数量专家总数即inTensor0Desc.shape.dims[1]要求可被dims[1]整除dims[1] % groupNum 0必传kint32_t0[1, groupNum]选择的 Top K 组数量必传groupMultiFlagGroupMultiFlag (uint16_t)UNDEFINED{UNDEFINED, SUM_MULTI_MAX}每组内取值计算方式UNDEFINED0为每组取最大值SUM_MULTI_MAX为每组内取 n 个最大值求和此时必须设置参数nnuint16_t1[1, expert_num / groupNum]每组内取值个数仅当groupMultiFlag SUM_MULTI_MAX时有效rsv[12]uint8_t[12]{0}—预留参数5.1 两种取值模式的计算差异groupMultiFlag UNDEFINED默认每组内取最大值作为该组得分等价于kInner 1。Runner 中注释明确 1: get the max value in each group。groupMultiFlag SUM_MULTI_MAX每组内取n个最大值求和作为组得分对应 MoE 路由中取每组前 n 个专家得分之和的聚合策略Runner 会将n直接映射为 Kernel 参数kInner。5.2 从源码看参数校验规则CreateOperationGroupTopkParamsrc/ops/ops_infer/group_topk/group_topk_operation.cpp#L29-L67在创建算子时执行第一道校验违反任一条件均返回ERROR_INVALID_PARAMgroupNum 1k 1且k groupNum保证选前 k 组有意义groupMultiFlag GroupMultiFlag::SUM_MULTI_MAX当groupMultiFlag SUM_MULTI_MAX时n 1平台限制仅当Config::Is910B()为真即 Atlas 800I A2 推理产品时允许创建其他平台直接报错。这一点与 Kernel 侧CanSupport、测试用例中的Ascend910BSocVersion 约束完全一致。InferShapeCheckImpl与SetupCheckImpl则进一步调用InTensorDescsChecksrc/ops/ops_infer/group_topk/group_topk_operation.cpp#L134-L195做张量形状维度的二次校验inTensor0token 得分矩阵必须为 2 维且dims[1]专家数取值[1, 1024]inTensor1索引数组必须为 1 维且dims[0] 1024与 Kernel 侧MAX_EXPERT常量一致groupNum需满足[1, expertNum]且expertNum % groupNum 0k需满足[1, groupNum]SUM_MULTI_MAX模式下n expertNum / groupNum。6. InferShape 与张量检查InferShapeImplsrc/ops/ops_infer/group_topk/group_topk_operation.cpp#L87-L92的实现非常简洁outTensorDescs.at(0) inTensorDescs.at(0)即输出形状与输入 0 完全一致。这是由算子的置零剪枝语义决定的——输出与输入同形只是非前 k 组位置的数值被清零。SetupCheckImpl除了再次调用InTensorDescsCheck外还通过TensorCheck::TensorDescsEqual校验输出 Tensor 的维数与各维与inTensor0一致outTensor0 dimNum/dims必须与inTensor0相同最后调用OutTensorsCheck确认输出为 2 维张量。7. Ops RunnerKernelGraph 构建与参数下传GroupTopkOpsRunner::SetupKernelGraphsrc/ops/ops_infer/group_topk/group_topk_ops_runner.cpp#L25-L63完成了 ATB 侧到 ASD 原生侧的最后一步组装图规模为 2 个输入、1 个输出、0 个内部张量、1 个节点tokenTensorkernelGraph_.inTensors[0]与idxArrTensorkernelGraph_.inTensors[1]分别对应两个输入构造AsdOps::OpParam::GroupTopk将 ATB 参数映射为原生参数groupNum、k直接透传kInner依据groupMultiFlag决定SUM_MULTI_MAX时取n否则取 1节点描述为{0, GroupTopkOperation, groupTopkParam}挂载输入输出张量指针。SetParamsrc/ops/ops_infer/group_topk/group_topk_ops_runner.cpp#L65-L73用于运行时参数热更新将新的GroupTopkParam与当前参数比较通过operator仅比较groupNum与k若不同则更新并置isParamUpdated_触发 KernelGraph 重建。文件末尾的REG_RUNNER_TYPE(GroupTopkOpsRunner)与REG_OP_PARAM(AsdOps::OpParam::GroupTopk)完成了 Runner 类型与原生参数类型的注册使框架能够按参数类型查找到对应的 Runner。8. Kernel 侧实现能力检查、Kernel 选择与 TilingKernel 侧位于 src/kernels/kernels/group_topk其职责是能不能跑、跑哪个 Kernel、怎么切分数据。8.1 Kernel 算子与 Kernel 选择src/kernels/kernels/group_topk/group_topk_operation.cpp 定义了 ASD 侧的GroupTopkOperation继承OperationBase其GetBestKernel中根据输入 dtype 选择 Kernel仅当inTensor0为TENSOR_DTYPE_FLOAT16或TENSOR_DTYPE_BF16且inTensor1为TENSOR_DTYPE_INT32时返回GroupTopkKernel否则返回空指针并报错。8.2 Kernel 能力检查GroupTopkKernel::CanSupportsrc/kernels/kernels/group_topk/group_topk_kernel.cpp#L33-L59对运行时输入做最终把关输入数量 2、输出数量 1inTensor0为 2 维[tokenNum, expertNum]tokenNum 00 expertNum 1024inTensor1为 1 维且长度必须等于 1024对应MAX_EXPERTgroupNum满足(0, expertNum]且expertNum % groupNum 0kInner满足(0, expertNum / groupNum]k满足(0, groupNum]。该文件顶部还声明了三个常量MAX_EXPERT 1024、MAX_GROUPNUM 256、MAX_KINNER 32构成算子的实现边界。GetTilingSize返回sizeof(GroupTopkTilingData)InitImpl调用GroupTopkTiling生成 Tiling 数据。8.3 Tiling 切分策略GroupTopkTilingsrc/kernels/kernels/group_topk/tiling/group_topk_tiling.cpp#L46-L82按 Vector Core 数量对 token 维度做数据切分读取PlatformInfo::Instance().GetCoreNum(CoreType::CORE_TYPE_VECTOR)得到核数计算expertNumPerGroup expertNum / groupNum并按对齐粒度RoundUp得到expertNumPerGroupPaddedkInner 1时按SORT_REPEAT_COUNT 32对齐否则按B16_PER_BLOCK对齐反映排序与单值两种路径的访存要求每个核平均分配tokenNum / coreNum个 token余数作为tailTokenNumkernelInfo.SetBlockDim(std::min(tokenNum, coreNum))确定实际启动的核数通过 tilingKey 区分不同变体kInner 1与否、BF16 与否、单值组特殊分支SINGLE_VALUE_GROUP_TILING_KEY供 Kernel 在编译期选择不同实现。9. 测试验证从 Kernel 单测到高精度用例仓库为该算子提供了多层测试可作为理解语义的活文档。9.1 Kernel 级单测tests/apitest/kernelstest/group_topk/test_group_topk.py 中实现了精确的 golden 计算逻辑可作为算子的参考伪代码将输入 reshape 为(token_num, group_num, expert_num // group_num)每组内取kInner个最大值并求和得到每组的组得分对组得分按降序argsort保留前k组将剩余组对应的位置置零最终输出与输入同形。测试通过generalize_param自动枚举 token/expert/groupNum/k/kInner 的组合覆盖expertNum的各类因子并在only_910b装饰器下分别跑 FP16 与 BF16另有两个固定用例test_fp16_512_160_8_3_2/test_bf16_512_160_8_3_2直接对应 DeepSeek-V3 的典型路由配置token512、expert256、groupNum8、k3、kInner2。9.2 高级功能测试tests/high_level_test/GroupTopkOperation/Requirements/GroupTopkOperation_DeepSeekV3_Precision.csvDeepSeek-V3 场景的精度用例同时覆盖参数合法性用例如groupNum161 expertNum报ERROR_INVALID_PARAM、groupMultiFlag-1/2非法值、非 910B 平台报错、n0报错等其SocVersion列统一为Ascend910B。tests/high_level_test/GroupTopkOperation/Smoke/GroupTopkOperation_TestCase.csv冒烟/泛化用例覆盖groupNum从 16 到 512、k从 16 到 128、shape 从10,1024到1024,1024的多种组合。tests/apitest/opstest/csv/group_topk.csvOps 级 CSV 用例用于框架统一驱动的算子测试流水线。这些测试共同验证了第 58 节所述的全部约束参数取值范围、整除要求、dtype 组合、平台限定仅 Atlas 800I A2 / Ascend910B以及置零语义的正确性。10. 快速导航本文路由文件原文.agent/knowledge/routing/group_topk.md详细知识条目.agent/knowledge/ops/other/group_topk/index.md知识库主索引.agent/knowledge/README.md算子源码src/ops/ops_infer/group_topk/Kernel 源码src/kernels/kernels/group_topk参数头文件include/atb/infer_op_params.h结语GroupTopkOperation虽然复杂度等级为 S简单但其代码组织完整覆盖了 ATB 算子框架的典型四层结构参数定义infer_op_params.h→ Operation 封装与校验group_topk_operation.→ OpsRunner 图构建与参数映射group_topk_ops_runner.→ Kernel 能力检查与 Tiling 切分src/kernels/kernels/group_topk。借助路由文件给出的阅读顺序开发者可以从头文件签名入手逐层追踪到原生 Kernel 的置零剪枝实现并通过三层测试用例验证任何参数组合的行为。对于希望在 ascend-transformer-boost 中新增或调试类似小而专推理算子的开发者本文梳理的路径与源码证据可以作为一份可直接复用的调研模板。【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

免费获取报价