CANN 自定义算子详解:npu_moe_gating_top_k 实现 MoE 路由与 TopK 专家选择
2026/9/19 23:41:34 网站建设 项目流程

CANN 自定义算子详解:npu_moe_gating_top_k 实现 MoE 路由与 TopK 专家选择

【免费下载链接】cann-recipes-infer本项目针对LLM与多模态模型推理业务中的典型模型、加速算法,提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-infer

本文围绕 CANN recipes-infer 仓库中 ops/ascendc 目录下的自定义算子npu_moe_gating_top_k展开,系统讲解其在 MoE(Mixture of Experts)推理中的核心作用:对 gating 分数完成 Sigmoid/Softmax/Softplus 归一化、分组排序、TopK 专家选择以及基于词表的 hash 路由,并给出完整参数说明、约束条件、调用示例与仓库源码级实现解析。读者阅读后可掌握该算子的数学语义、入参约束、PyTorch 调用方式,以及它从 Host 侧 Tiling 到 NPU Kernel 的完整实现路径,可直接在 Atlas A3 推理系列与 Ascend 950 系列产品上落地使用。

产品支持情况

该算子已适配以下昇腾产品形态,可直接用于推理场景:

产品是否支持
Atlas A3 推理系列产品
Ascend 950PR / Ascend 950DT

从算子定义源码 moe_gating_top_k_hash_def.cpp 可以看出,MoeGatingTopKHash通过AICore().AddConfig("ascend910b")AddConfig("ascend910_93")以及为ascend950配置独立的regbaseCfg(开启动态编译、动态 rank 与动态 shape 支持)来声明各平台的调度配置,其中ascend950走寄存器基址(regbase)专用 kernel 分支。

功能说明:MoE 路由的一站式融合算子

在 MoE(混合专家)模型中,路由网络(gating network)决定每个 token 应被送往哪些专家(expert)处理。npu_moe_gating_top_k将 gating 计算中的多个步骤融合为单个算子,避免多次 kernel 启动与中间张量的反复搬运,整体流程如下:

  1. 对输入 gating 分数x做归一化(Sigmoid / Softmax / Softplus),可选叠加偏置;
  2. group_count对结果分组,每组内部先做 TopK,按组分数(最大值或 top2 之和)排序并选出前k_group个组;
  3. 若提供input_idstid2eid,则直接依据词表映射完成 hash 路由得到专家索引;否则在选中的组内再做 TopK,得到最终专家索引;
  4. 对选出的分数按routed_scaling_factoreps做缩放归一化,得到最终的专家路由权重。

整个算子一次前向即可输出归一化分数、专家索引以及(可选的)归一化中间结果,供后续专家计算(如 Grouped GEMM / GMM)直接消费。

归一化公式

对输入xnorm_type选择归一化方式:

$$ \begin{aligned} &if\ normType == 1: normOut=Sigmoid(x) \ &else\ if\ normType == 0: normOut=SoftMax(x) \ &else\ : normOut=Softplus(x) \end{aligned} $$

bias不为空:

$$ normOut = normOut + bias $$

additional_biasadditional_token_mask同时提供,对于additional_token_mask中取值为 true 的行,使用additional_bias替换bias参与上述计算:

$$ normOut[row] = normOut[row] + additional_bias,\ where\ additional_token_mask[row] == True $$

分组排序公式

group_count对计算结果分组,每组按group_select_mode取 max 或 topk2 的 sum 值对组进行排序,取前k_group个组:

$$ groupOut, groupId = TopK(ReduceSum(TopK(Split(normOut, groupCount), k=2, dim=-1), dim=-1),k=kGroup) $$

专家选择公式

  • 若指定了input_idstid2eid,则根据输入的词表进行 hash 操作;
  • 否则根据上一步的groupId获取normOut中对应的元素,将数据再做 TopK,得到expertIdxOut

$$ y,expertIdxOut=TopK(normOut[groupId, :],k=k) $$

路由权重缩放公式

y按照输入的routedScalingFactoreps参数计算,得到yOut

$$ yOut = y / (ReduceSum(y, dim=-1)+eps)*routedScalingFactor $$

函数原型

custom.npu_moe_gating_top_k(Tensor x, int k, *, Tensor? bias=None, Tensor? input_ids=None, Tensor? tid2eid=None, Tensor? additional_bias=None, Tensor? additional_token_mask=None, int k_group=1, int group_count=1, float routed_scaling_factor=1., float eps=9.9999999999999995e-21, int group_select_mode=0, int renorm=0, int norm_type=0, bool out_flag=False) -> (Tensor, Tensor, Tensor)

说明:b(batch size)表示输入样本批量大小、s(sequence length)表示输入样本序列长度、T 表示 bs 合轴后的大小、e 表示专家数量、k 表示选取 Top 专家数。

参数说明

参数名类型描述
xTensor输入张量,支持 2D 或 3D,shape 为 (T, e) 或 (b, s, e),支持float16bfloat16float32
kint选取的专家数量,取值小于等于 e,且必须小于等于 64
biasTensor(可选)偏置张量,shape 为(e),dtype 与 x 相同,支持float16bfloat16float32
input_idsTensor(可选)输入词表,shape 为(T),仅支持int64,取值范围为 [0, n],n 为 tid2eid 第一维的大小
tid2eidTensor(可选)词表到专家 id 的映射关系表,shape 为(n,k),仅支持int32,取值范围为 [0, e],e 代表专家数
additional_biasTensor(可选)附加偏置张量,shape 为(e),dtype 与 x 相同。仅在 additional_token_mask 同时提供时生效:additional_token_mask 为 true 的行使用 additional_bias 替换 bias 参与 gating 计算
additional_token_maskTensor(可选)附加 token 标记,shape 为(T),仅支持bool,取值为 true 表示该 token 使用 additional_bias 替换 bias
k_groupint(可选)选取的组数量,默认为 1
group_countint(可选)总组数,默认为 1
routed_scaling_factorfloat(可选)路由缩放因子,默认为 1
epsfloat(可选)数值稳定性参数,防止除零,默认为 1e-20
group_select_modeint(可选)组选择模式:0-使用最大值排序,1-使用 top2 的和排序
renormint(可选)重归一化标志,仅支持 0
norm_typeint(可选)归一化标志,0-Softmax, 1-Sigmoid, 2-Softplus
out_flagbool(可选)是否输出归一化结果

其中默认属性值与算子 Host 侧定义一致:在 moe_gating_top_k_hash_def.cpp 中,k_groupgroup_countgroup_select_moderenormnorm_type的默认值分别为 1、1、0、0、0,out_flag默认 false,routed_scaling_factor默认 1.0,eps默认 1e-20f。

返回值说明

返回值类型描述
yTensor归一化、分组排序和 TopK 后的结果
expert_idxTensor专家索引,数据类型为 int32
outTensor归一化结果(当 out_flag=True 时有效)

在 npu_moe_gating_top_k.cpp 中,输出张量的 shape 推导逻辑为:yOutexpertIdxOut保持输入x除最后一维外的所有维度,最后一维替换为k;其中expertIdxOut固定为int32normOut与输入xshape 完全一致,但统一使用float32存储归一化中间结果。该实现使用sym_size/empty_symint保留动态维度(如 token 维),避免动态 shape 场景下符号维被过早具象化。

约束说明

使用该算子时需严格遵守以下约束:

  • renorm仅支持 0,表示先进行 norm 操作,再计算 topk。
  • group_select_mode取值 0 和 1,0 表示使用最大值对 group 进行排序,1 表示使用 topk2 的 sum 值对 group 排序。
  • norm_type取值 0、1 和 2,0 表示使用 Softmax 函数,1 表示使用 Sigmoid 函数,2 表示使用 Softplus 函数。
  • out_flag取值 true 和 false,true 表示输出,false 表示不输出。
  • input_idstid2eid都不为空表示 hash 场景,都为空表示 topk 场景,不允许只有一个为空。
  • k_groupgroup_count为 1 时,表示不分组排序。
  • bias的 dtype 要和 x 相同。
  • additional_bias的 dtype 要和 x 相同,shape 为(e);additional_token_mask仅支持 bool,shape 为(T)。
  • additional_bias仅在additional_token_mask同时提供时生效,additional_token_mask为 true 的行使用additional_bias替换bias
  • 该接口支持推理场景下使用。
  • 该接口支持 aclgraph 入图。
  • 该接口与 PyTorch 配合使用时,需要保证 CANN 相关包与 PyTorch 相关包的版本匹配。

上述大部分约束在 Torch 扩展的入口处也有显式校验(TORCH_CHECK),例如 npu_moe_gating_top_k.cpp 会校验k > 0kGroup > 0groupCount > 0k <= x_shape[-1] / groupCount * kGroupkGroup <= groupCountgroupSelectMode仅为 0 或 1、normType仅为 0/1/2、renorm仅为 0;同时校验bias/additional_bias为一维且长度等于x最后一维、additional_token_mask为 bool 且长度等于x的行数。若传参不合法,会在 Host 侧直接报错,方便快速定位问题。

调用示例

仓库提供了完整的可运行示例与单元测试:test_npu_moe_gating_top_k.py。该脚本依赖torchtorch_nputorchaircustom_ops(即本仓库 torch_ops_extension 编译出的自定义算子包),并在固定随机种子下对比 NPU 结果与 NumPy CPU 参考实现,验证精度。

基础 TopK 场景调用

以下代码展示了最常见的 TopK 场景(不分组、无 hash),对应测试用例test_moe_gating_top_k_384_experts_topk6

import torch import torch_npu import custom_ops DEVICE_ID = 0 torch_npu.npu.set_device(int(DEVICE_ID)) batch_size = 16 expert_count = 384 k = 6 k_group = 1 group_count = 1 routed_scaling_factor = 1.0 eps = 1e-6 group_select_mode = 0 renorm = 0 norm_type = 1 # Sigmoid out_flag = False x = torch.randn(batch_size, expert_count, dtype=torch.float16).npu() y_out, expert_idx, _ = torch.ops.custom.npu_moe_gating_top_k( x, k, bias=None, input_ids=None, tid2eid=None, k_group=k_group, group_count=group_count, routed_scaling_factor=routed_scaling_factor, eps=eps, group_select_mode=group_select_mode, renorm=renorm, norm_type=norm_type, out_flag=out_flag ) # y_out: (batch_size, k),专家路由权重 # expert_idx: (batch_size, k),int32 专家索引

Hash 路由场景调用

当提供input_idstid2eid时,专家索引不再依赖 TopK 计算,而是直接查表完成 hash 路由,对应测试用例test_moe_gating_top_k_different_dtypes

N = 100 # 词表大小 input_ids = torch.randint(0, N, (batch_size,), dtype=torch.int64).npu() tid2eid = torch.randint(0, expert_count, (N, k), dtype=torch.int32).npu() y_out, expert_idx, _ = torch.ops.custom.npu_moe_gating_top_k( x, k, bias=None, input_ids=input_ids, tid2eid=tid2eid, k_group=1, group_count=1, routed_scaling_factor=1.0, eps=1e-6, group_select_mode=0, renorm=0, norm_type=1, out_flag=False )

additional_bias 场景调用

测试用例test_moe_gating_top_k_additional_bias展示了按 token 粒度替换偏置的用法:一半 token 标记为使用additional_bias,其余仍使用普通bias,覆盖 softmax / sigmoid / softplus 三种归一化模式:

bias = torch.rand(expert_count, dtype=torch.float16).npu() additional_bias = torch.rand(expert_count, dtype=torch.float16).npu() additional_token_mask = torch.zeros(batch_size, dtype=torch.bool) additional_token_mask[::2] = True # 一半 token 使用 additional_bias additional_token_mask = additional_token_mask.npu() y_out, expert_idx, _ = torch.ops.custom.npu_moe_gating_top_k( x, k, bias=bias, input_ids=None, tid2eid=None, additional_bias=additional_bias, additional_token_mask=additional_token_mask, k_group=1, group_count=1, routed_scaling_factor=1.0, eps=1e-6, group_select_mode=0, renorm=0, norm_type=0, out_flag=False )

torch.compile 图模式调用

同一测试文件中的test_moe_gating_top_k_different_dtypes_graph展示了通过torch.compile+torchairNPU 后端将算子接入计算图(aclgraph 入图)的方式:

from torchair.configs.compiler_config import CompilerConfig class Network(torch.nn.Module): def forward(self, x_npu, k, bias, input_ids, tid2eid, k_group, group_count, routed_scaling_factor, eps, group_select_mode, renorm, norm_type, out_flag): y_out, expert_idx, y2_out = torch.ops.custom.npu_moe_gating_top_k( x_npu, k, bias=bias, input_ids=input_ids, tid2eid=tid2eid, k_group=k_group, group_count=group_count, routed_scaling_factor=routed_scaling_factor, eps=eps, group_select_mode=group_select_mode, renorm=renorm, norm_type=norm_type, out_flag=out_flag ) return y_out, expert_idx, y2_out config = CompilerConfig() config.mode = "reduce-overhead" npu_backend = torchair.get_npu_backend(compiler_config=config) model = torch.compile(Network().npu(), fullgraph=True, backend=npu_backend, dynamic=False) y_out, expert_idx, _ = model(x_npu, k, None, input_ids_npu, tid2eid_npu, k_group, group_count, routed_scaling_factor, eps, group_select_mode, renorm, norm_type, out_flag)

源码实现剖析

Torch 侧算子注册与图转换

自定义算子通过 ops_def_registration.cpp 注册到torch.ops.custom命名空间,并在 npu_moe_gating_top_k.cpp 中分别注册PrivateUse1(NPU 设备)与Meta(shape 推导)两套实现。

在 torch.compile 场景下,npu_moe_gating_top_k.py 通过@register_fx_node_ge_converter(torch.ops.custom.npu_moe_gating_top_k.default)将 FX 节点转换为 GE 自定义算子节点MoeGatingTopKHash,其中meta_outputs形参为固定写法,用于推导 GE 节点的输出 dtype 与 shape;算子的 6 个输入映射到inputs,9 个标量参数(k、k_group、group_count、routed_scaling_factor、eps、group_select_mode、renorm、norm_type、out_flag)映射为attrs,输出为yexpert_idxout

Host 侧 Tiling 与 Workspace

在 moe_gating_top_k_hash_tiling.h 中定义了MoeGatingTopKHashTilingData,包含needCoreNumrowCountperCoreRowCountlastCoreRowCountexpertCountaddBiaskkGroupgroupCountperGroupExpertCountgroupSelectModerenormnormTypeoutFlaghashFlagroutedScalingFactoreps以及 Softmax Tiling 结构体等字段,用于把 shape 与属性翻译成 Kernel 可消费的切分参数。

在 moe_gating_top_k_hash_tiling.cpp 中,Tiling 流程包含获取平台资源(CoreNum、UB/L1/L0C 大小)、读取输入输出与属性、按行数切分任务、计算 TilingKey、申请 Workspace(默认预留 16M,见DEFAULT_WORKSPACE_SIZE)等步骤。可以看到算子在 Host 侧还定义了多种 TilingKey 分支:专家数/组数对齐的高性能分支、不分组分支(WITHOUT_GROUP)、通用分支(GENERALIZED)等,其中不分组场景又按input_ids/tid2eid的 int32/int64 组合细分为多个模板实例。

Kernel 侧多分支调度

在 moe_gating_top_k_hash.cpp 中,Kernel 入口根据 TilingKey 分发到不同实现类:MoeGatingTopKHashEKFullload(每组专家数对齐高性能路径)、MoeGatingTopKHashWithoutGroup(不分组场景,含 int32/int64 索引组合的多个模板实例)与MoeGatingTopKHashGenerlized(通用分组场景),并在__DAV_C310__编译条件下额外引入MoeGatingTopKHashRegbase寄存器基址实现,对应 Ascend 950 平台的regbaseCfg配置。Kernel 声明为KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY),仅使用向量核执行,且对 AIC 核直接返回,确保任务只落在 AIV 上。

与 MoE 推理流水线的衔接

该算子的输出(路由权重y与专家索引expert_idx)是 MoE 前向中专家并行(EP)与分组 GEMM 的前置输入。在 CANN 推理样例仓库中,它常与 MoE 相关的其他自定义算子(如npu_moe_init_routing_group_quantnpu_swiglu_group_quantnpu_moe_*系列,见 ops/ascendc/docs)组合使用:先由本算子完成路由决策,再由后续量化与矩阵乘算子按expert_idx聚合 token、完成各专家的前向计算。将 gating 的归一化、偏置、分组排序、TopK/hash、权重缩放融合为单算子,可显著减少中间张量的落盘与多次 kernel 启动开销,是 MoE 推理优化的常用手段之一。

使用注意事项

  1. 本算子面向推理场景,训练场景请另行评估;
  2. 与 PyTorch 配合使用时,务必保证 CANN 相关包(torch_npu、torchair)与 PyTorch 版本匹配,否则可能出现接口或算子注册不兼容的问题;
  3. hash 场景中input_idstid2eid必须同时提供或同时为空,二者取值的上下界([0, n] 与 [0, e])需要由调用方保证,越界会导致非法专家索引;
  4. k上限为 64,且不能超过x_shape[-1] / group_count * k_group,请结合模型实际专家数与分组策略设置;
  5. 精度验证可直接复用 test_npu_moe_gating_top_k.py 中的 NumPy 参考实现moe_gating_top_k_numpy,其覆盖 softmax/sigmoid/softplus、bias 与 additional_bias 等多种组合,可作为自测基线。

【免费下载链接】cann-recipes-infer本项目针对LLM与多模态模型推理业务中的典型模型、加速算法,提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-infer

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询