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 启动与中间张量的反复搬运,整体流程如下:
- 对输入 gating 分数
x做归一化(Sigmoid / Softmax / Softplus),可选叠加偏置; - 按
group_count对结果分组,每组内部先做 TopK,按组分数(最大值或 top2 之和)排序并选出前k_group个组; - 若提供
input_ids与tid2eid,则直接依据词表映射完成 hash 路由得到专家索引;否则在选中的组内再做 TopK,得到最终专家索引; - 对选出的分数按
routed_scaling_factor与eps做缩放归一化,得到最终的专家路由权重。
整个算子一次前向即可输出归一化分数、专家索引以及(可选的)归一化中间结果,供后续专家计算(如 Grouped GEMM / GMM)直接消费。
归一化公式
对输入x按norm_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_bias与additional_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_ids和tid2eid,则根据输入的词表进行 hash 操作; - 否则根据上一步的
groupId获取normOut中对应的元素,将数据再做 TopK,得到expertIdxOut:
$$ y,expertIdxOut=TopK(normOut[groupId, :],k=k) $$
路由权重缩放公式
对y按照输入的routedScalingFactor和eps参数计算,得到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 专家数。
参数说明
| 参数名 | 类型 | 描述 |
|---|---|---|
x | Tensor | 输入张量,支持 2D 或 3D,shape 为 (T, e) 或 (b, s, e),支持float16、bfloat16和float32 |
k | int | 选取的专家数量,取值小于等于 e,且必须小于等于 64 |
bias | Tensor(可选) | 偏置张量,shape 为(e),dtype 与 x 相同,支持float16、bfloat16和float32 |
input_ids | Tensor(可选) | 输入词表,shape 为(T),仅支持int64,取值范围为 [0, n],n 为 tid2eid 第一维的大小 |
tid2eid | Tensor(可选) | 词表到专家 id 的映射关系表,shape 为(n,k),仅支持int32,取值范围为 [0, e],e 代表专家数 |
additional_bias | Tensor(可选) | 附加偏置张量,shape 为(e),dtype 与 x 相同。仅在 additional_token_mask 同时提供时生效:additional_token_mask 为 true 的行使用 additional_bias 替换 bias 参与 gating 计算 |
additional_token_mask | Tensor(可选) | 附加 token 标记,shape 为(T),仅支持bool,取值为 true 表示该 token 使用 additional_bias 替换 bias |
k_group | int(可选) | 选取的组数量,默认为 1 |
group_count | int(可选) | 总组数,默认为 1 |
routed_scaling_factor | float(可选) | 路由缩放因子,默认为 1 |
eps | float(可选) | 数值稳定性参数,防止除零,默认为 1e-20 |
group_select_mode | int(可选) | 组选择模式:0-使用最大值排序,1-使用 top2 的和排序 |
renorm | int(可选) | 重归一化标志,仅支持 0 |
norm_type | int(可选) | 归一化标志,0-Softmax, 1-Sigmoid, 2-Softplus |
out_flag | bool(可选) | 是否输出归一化结果 |
其中默认属性值与算子 Host 侧定义一致:在 moe_gating_top_k_hash_def.cpp 中,k_group、group_count、group_select_mode、renorm、norm_type的默认值分别为 1、1、0、0、0,out_flag默认 false,routed_scaling_factor默认 1.0,eps默认 1e-20f。
返回值说明
| 返回值 | 类型 | 描述 |
|---|---|---|
y | Tensor | 归一化、分组排序和 TopK 后的结果 |
expert_idx | Tensor | 专家索引,数据类型为 int32 |
out | Tensor | 归一化结果(当 out_flag=True 时有效) |
在 npu_moe_gating_top_k.cpp 中,输出张量的 shape 推导逻辑为:yOut与expertIdxOut保持输入x除最后一维外的所有维度,最后一维替换为k;其中expertIdxOut固定为int32。normOut与输入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_ids和tid2eid都不为空表示 hash 场景,都为空表示 topk 场景,不允许只有一个为空。k_group和group_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 > 0、kGroup > 0、groupCount > 0、k <= x_shape[-1] / groupCount * kGroup、kGroup <= groupCount、groupSelectMode仅为 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。该脚本依赖torch、torch_npu、torchair与custom_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_ids与tid2eid时,专家索引不再依赖 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,输出为y、expert_idx、out。
Host 侧 Tiling 与 Workspace
在 moe_gating_top_k_hash_tiling.h 中定义了MoeGatingTopKHashTilingData,包含needCoreNum、rowCount、perCoreRowCount、lastCoreRowCount、expertCount、addBias、k、kGroup、groupCount、perGroupExpertCount、groupSelectMode、renorm、normType、outFlag、hashFlag、routedScalingFactor、eps以及 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_quant、npu_swiglu_group_quant、npu_moe_*系列,见 ops/ascendc/docs)组合使用:先由本算子完成路由决策,再由后续量化与矩阵乘算子按expert_idx聚合 token、完成各专家的前向计算。将 gating 的归一化、偏置、分组排序、TopK/hash、权重缩放融合为单算子,可显著减少中间张量的落盘与多次 kernel 启动开销,是 MoE 推理优化的常用手段之一。
使用注意事项
- 本算子面向推理场景,训练场景请另行评估;
- 与 PyTorch 配合使用时,务必保证 CANN 相关包(torch_npu、torchair)与 PyTorch 版本匹配,否则可能出现接口或算子注册不兼容的问题;
- hash 场景中
input_ids与tid2eid必须同时提供或同时为空,二者取值的上下界([0, n] 与 [0, e])需要由调用方保证,越界会导致非法专家索引; k上限为 64,且不能超过x_shape[-1] / group_count * k_group,请结合模型实际专家数与分组策略设置;- 精度验证可直接复用 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),仅供参考