ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

CANN 自定义算子详解:npu_moe_gating_top_k 实现 MoE 路由与 TopK 专家选择 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展开系统讲解其在 MoEMixture 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: normOutSigmoid(x) \ else\ if\ normType 0: normOutSoftMax(x) \ else\ : normOutSoftplus(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), k2, dim-1), dim-1),kkGroup) $$专家选择公式若指定了input_ids和tid2eid则根据输入的词表进行 hash 操作否则根据上一步的groupId获取normOut中对应的元素将数据再做 TopK得到expertIdxOut$$ y,expertIdxOutTopK(normOut[groupId, :],kk) $$路由权重缩放公式对y按照输入的routedScalingFactor和eps参数计算得到yOut$$ yOut y / (ReduceSum(y, dim-1)eps)*routedScalingFactor $$函数原型custom.npu_moe_gating_top_k(Tensor x, int k, *, Tensor? biasNone, Tensor? input_idsNone, Tensor? tid2eidNone, Tensor? additional_biasNone, Tensor? additional_token_maskNone, int k_group1, int group_count1, float routed_scaling_factor1., float eps9.9999999999999995e-21, int group_select_mode0, int renorm0, int norm_type0, bool out_flagFalse) - (Tensor, Tensor, Tensor)说明bbatch size表示输入样本批量大小、ssequence length表示输入样本序列长度、T 表示 bs 合轴后的大小、e 表示专家数量、k 表示选取 Top 专家数。参数说明参数名类型描述xTensor输入张量支持 2D 或 3Dshape 为 (T, e) 或 (b, s, e)支持float16、bfloat16和float32kint选取的专家数量取值小于等于 e且必须小于等于 64biasTensor可选偏置张量shape 为edtype 与 x 相同支持float16、bfloat16和float32input_idsTensor可选输入词表shape 为T仅支持int64取值范围为 [0, n]n 为 tid2eid 第一维的大小tid2eidTensor可选词表到专家 id 的映射关系表shape 为nk仅支持int32取值范围为 [0, e]e 代表专家数additional_biasTensor可选附加偏置张量shape 为edtype 与 x 相同。仅在 additional_token_mask 同时提供时生效additional_token_mask 为 true 的行使用 additional_bias 替换 bias 参与 gating 计算additional_token_maskTensor可选附加 token 标记shape 为T仅支持bool取值为 true 表示该 token 使用 additional_bias 替换 biask_groupint可选选取的组数量默认为 1group_countint可选总组数默认为 1routed_scaling_factorfloat可选路由缩放因子默认为 1epsfloat可选数值稳定性参数防止除零默认为 1e-20group_select_modeint可选组选择模式0-使用最大值排序1-使用 top2 的和排序renormint可选重归一化标志仅支持 0norm_typeint可选归一化标志0-Softmax, 1-Sigmoid, 2-Softplusout_flagbool可选是否输出归一化结果其中默认属性值与算子 Host 侧定义一致在 moe_gating_top_k_hash_def.cpp 中k_group、group_count、group_select_mode、renorm、norm_type的默认值分别为 1、1、0、0、0out_flag默认 falserouted_scaling_factor默认 1.0eps默认 1e-20f。返回值说明返回值类型描述yTensor归一化、分组排序和 TopK 后的结果expert_idxTensor专家索引数据类型为 int32outTensor归一化结果当 out_flagTrue 时有效在 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 和 10 表示使用最大值对 group 进行排序1 表示使用 topk2 的 sum 值对 group 排序。norm_type取值 0、1 和 20 表示使用 Softmax 函数1 表示使用 Sigmoid 函数2 表示使用 Softplus 函数。out_flag取值 true 和 falsetrue 表示输出false 表示不输出。input_ids和tid2eid都不为空表示 hash 场景都为空表示 topk 场景不允许只有一个为空。k_group和group_count为 1 时表示不分组排序。bias的 dtype 要和 x 相同。additional_bias的 dtype 要和 x 相同shape 为eadditional_token_mask仅支持 boolshape 为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_topk6import 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, dtypetorch.float16).npu() y_out, expert_idx, _ torch.ops.custom.npu_moe_gating_top_k( x, k, biasNone, input_idsNone, tid2eidNone, k_groupk_group, group_countgroup_count, routed_scaling_factorrouted_scaling_factor, epseps, group_select_modegroup_select_mode, renormrenorm, norm_typenorm_type, out_flagout_flag ) # y_out: (batch_size, k)专家路由权重 # expert_idx: (batch_size, k)int32 专家索引Hash 路由场景调用当提供input_ids与tid2eid时专家索引不再依赖 TopK 计算而是直接查表完成 hash 路由对应测试用例test_moe_gating_top_k_different_dtypesN 100 # 词表大小 input_ids torch.randint(0, N, (batch_size,), dtypetorch.int64).npu() tid2eid torch.randint(0, expert_count, (N, k), dtypetorch.int32).npu() y_out, expert_idx, _ torch.ops.custom.npu_moe_gating_top_k( x, k, biasNone, input_idsinput_ids, tid2eidtid2eid, k_group1, group_count1, routed_scaling_factor1.0, eps1e-6, group_select_mode0, renorm0, norm_type1, out_flagFalse )additional_bias 场景调用测试用例test_moe_gating_top_k_additional_bias展示了按 token 粒度替换偏置的用法一半 token 标记为使用additional_bias其余仍使用普通bias覆盖 softmax / sigmoid / softplus 三种归一化模式bias torch.rand(expert_count, dtypetorch.float16).npu() additional_bias torch.rand(expert_count, dtypetorch.float16).npu() additional_token_mask torch.zeros(batch_size, dtypetorch.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, biasbias, input_idsNone, tid2eidNone, additional_biasadditional_bias, additional_token_maskadditional_token_mask, k_group1, group_count1, routed_scaling_factor1.0, eps1e-6, group_select_mode0, renorm0, norm_type0, out_flagFalse )torch.compile 图模式调用同一测试文件中的test_moe_gating_top_k_different_dtypes_graph展示了通过torch.compiletorchairNPU 后端将算子接入计算图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, biasbias, input_idsinput_ids, tid2eidtid2eid, k_groupk_group, group_countgroup_count, routed_scaling_factorrouted_scaling_factor, epseps, group_select_modegroup_select_mode, renormrenorm, norm_typenorm_type, out_flagout_flag ) return y_out, expert_idx, y2_out config CompilerConfig() config.mode reduce-overhead npu_backend torchair.get_npu_backend(compiler_configconfig) model torch.compile(Network().npu(), fullgraphTrue, backendnpu_backend, dynamicFalse) 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 中分别注册PrivateUse1NPU 设备与Metashape 推导两套实现。在 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 个输入映射到inputs9 个标量参数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),仅供参考
返回列表