FEATURED · 精选文章

基于 PyPTO 的 MoE 融合算子实践:grouped_matmul_finalize_routing 的 MXFP8 实现与精度验证

发布时间 / 2026/9/18 18:45:42
来源 / 创域科博编辑部
栏目 / 资讯中心
基于 PyPTO 的 MoE 融合算子实践:grouped_matmul_finalize_routing 的 MXFP8 实现与精度验证 基于 PyPTO 的 MoE 融合算子实践grouped_matmul_finalize_routing 的 MXFP8 实现与精度验证【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gymgrouped_matmul_finalize_routing 是 CANN pypto-gym 仓库中面向 MoEMixture of Experts推理场景的 grouped matmul 后处理融合算子对应aclnnGroupedMatmulFinalizeRoutingV3的 MXFP8 路径。本文以 关联文档 为核心骨架结合 kernel 实现、golden 参考实现 与 单测入口 展开讲清该算子的语义、输入输出规格、Shape 约束、kernel 实现细节与验证方法帮助读者在 PyPTO 框架下快速理解并复现这类分组矩阵乘 路由后处理融合算子。一、产品支持情况关联文档明确标注了该算子当前的平台支持范围即当前仓库中的验证结论不代表未来版本Ascend 950PR不支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持需要说明的是其上级目录 matmul 算子目录说明 整体标注为 Ascend 950PR / Atlas A3 / Atlas A2 支持而本算子的单测在 test_gmm_finalize_routing.py 中带有pytest.mark.soc(950)标记测试标记与文档标注存在差异。实际落地时请以目标硬件上运行单测的结果为准本文描述的实现细节与数值行为以仓库当前代码为准。二、算子语义与数学公式grouped_matmul_finalize_routing是 MoE 场景中的 grouped matmul 后处理融合算子。它把按 expert 分组做矩阵乘与路由结果回写两件事融合在一个 kernel 中完成具体包含三部分工作将路由后的 token 按 expert 分组执行矩阵乘MXFP8 scaled matmul输出 FP32用每个 token 的logit对 matmul 结果做加权按row_index将加权结果 scatter-add 回最终输出并叠加 shared expert 的输出。2.1 数学公式mm_i ScaledMatmul(x1_i, x2_i, pertoken_scale_i, scale) weighted_i mm_i * logit_i out[row_index_i] weighted_i out[shared_input_offset:shared_input_offsetbatch] shared_input * shared_input_weight展开形式为out[row_index[t], n] logit[t] * Σ(k0..K-1) dequant(x1[t, k]) * dequant(x2[expert(t), k, n])其中dequant由 MXFP8 输入值和 E8M0FNU scale 共同决定即块级缩放反量化。关于 MXFP8每 64 个元素共享一个缩放因子缩放因子为仅含指数部分的 E8M0FNU 格式数据部分为 8 位浮点E4M3FN / E5M2详见 matmul 目录说明 中对 MXFP8 的注释。2.2 计算流程整个算子按以下四个阶段执行阶段 1Grouped Matmul 计算。按 expert 切分 token调用pypto.scaled_mm计算 FP32 输出x1_i: [M_i, K]x2_i: [K, N]或[N, K]由transpose_x2决定输出mm_i: [M_i, N]阶段 2Logit 加权。当has_logitTrue时对每个 token 的 matmul 结果乘以对应logitlogit_i: [M_i]→ unsqueeze →[M_i, 1]广播乘法[M_i, N] × [M_i, 1] → [M_i, N]阶段 3Finalize Routing 回写。根据row_index将 expert 输出累加到最终输出row_index_i: [M_i]out[row_index_i] weighted_i阶段 4Shared Expert 叠加。当has_shared_inputTrue时将 shared expert 输出按权重叠加到outshared_input: [batch, N]out[offset:offsetbatch] shared_input * shared_input_weight三、输入输出规格3.1 输入张量名称ShapeDType说明x1[M, K]FP8 E4M3/E5M2路由后 token 输入x2[E, K, N]或[E, N, K]FP8 E4M3/E5M2expert 权重布局由transpose_x2决定scale[ceil(K/64), N, 2]或[N, ceil(K/64), 2]E8M0FNU权重 scalepertoken_scale[M, ceil(K/64), 2]E8M0FNUtoken scalegroup_list[E]int64expert 分组信息shared_input[batch, N]bfloat16shared expert 输出logit[M]float32token 对应 expert 权重row_index[M]int64输出行索引out[batch, N]float32输出初值3.2 输出张量名称ShapeDType说明output[batch, N]float32finalize routing 后的融合结果3.3 源码中的张量构造在 test_gmm_finalize_routing.py 的_build_finalize_routing_tensors中可以看到与上述规格一致的构造方式x1 torch.randn((config.m, config.k), ...).to(torch_dtype)其中torch_dtype按in_dtype映射为torch.float8_e4m3fn或torch.float8_e5m2transpose_x2True时x2形状为[E, N, K]scale形状为[N, ceil(K/64), 2]否则为[E, K, N]与[ceil(K/64), N, 2]scale与pertoken_scale均以torch.float8_e8m0fnu存储K 维分块数scale_k (k 63) // 64row_index torch.arange(config.m) % config.batch保证索引落在[0, batch)shared_input使用 bfloat16out为 FP32 全零初值。四、Shape 范围与约束4.1 动态轴当前覆盖范围轴当前覆盖范围说明batch{64, 128, 256}输出行数M{128, 256, 768}路由后 token 数K{5120, 6144, 7168, 8192}matmul K 维N4096输出列数E{8, 16, 32}expert 数量4.2 约束条件transpose_x1 仅支持 False当前目标路径不支持转置x1。这一点在 golden 中也被强制校验——gen_golden中若cfg.transpose_x1为 True 会直接raise ValueError(aclnnGroupedMatmulFinalizeRoutingV3 only supports transposeX1False.)。group_list 当前测试为均匀分组kernel 内按M // E切分 token。row_index 范围合法row_index中元素必须位于[0, batch)。shared_input 边界合法shared_input_offset shared_input.shape[0] out.shape[0]。MXFP8 scale 布局固定K 维按ceil(K/64)分块每个 block 包含 2 个 E8M0FNU scale。4.3 group_list 的两种格式从 golden 实现 的_expert_range可以看出group_list_type支持两种分组描述格式配置项见FinalizeRoutingConfig.group_list_type当前测试取值为 1group_list_type 0前缀和格式第i个 expert 的 token 区间为[group_list[i-1], group_list[i])group_list_type 1计数格式每个元素为对应 expert 的 token 数第i个 expert 的区间为[sum(group_list[:i]), sum(group_list[:i]) group_list[i])。测试工具函数make_group_list支持生成这两种格式并允许M不能被E整除时把余量分摊到前面的 expert。五、PyPTO Kernel 实现解析核心实现位于 gmm_finalize_routing_impl.py由三部分组成配置数据结构FinalizeRoutingConfig、JIT kernelgmm_finalize_routing_kernel、host 侧封装gen_pypto。5.1 配置结构 FinalizeRoutingConfig配置字段覆盖了算子的全部行为开关batch输出 batch 维大小shared_input 的行数基准topk每个 batch 的 token 数mtoken 总数由batch * topk自动计算不可手动指定k/nmatmul 的 K 维与 N 维输出列数num_expertsexpert 数量in_dtype输入 FP8 数据类型默认pypto.DT_FP8E4M3transpose_x1/transpose_x2是否转置输入transpose_x1当前仅支持 Falsegroup_list_type0前缀和1每组计数shared_input_weight默认 1.0与shared_input_offset默认 0shared_input 叠加权重与起始行偏移has_logit默认 True与has_shared_input默认 True两个后处理分支开关vector_tile_shape向量算子 tile 配置。值得关注的是__post_init__中的自动 tile 推导逻辑kernel 内每个 expert 分到的 token 数为per_expert_m m // num_experts根据其大小自动选择 cube tileper_expert_m 64m_tile_shape[per_expert_m, per_expert_m]k_tile_shape[256, 512]n_tile_shape[256, 512]per_expert_m 1024m_tile_shape[128, 128]k_tile_shape[512, 512]n_tile_shape[128, 256]其余情况m_tile_shape[128, 128]k_tile_shape[256, 256]n_tile_shape[256, 256]。这保证了 tile 配置随 Shape 自动适配无需手写。5.2 JIT 编译选项kernel 通过pypto.frontend.jit装饰器编译带有两组关键选项pass_options{ cube_nbuffer_setting: {-1: 1}, vec_nbuffer_setting: {-2: 1, -1: 1}, auto_mix_partition: 1, }, runtime_options{ stitch_function_max_num: 128, device_sched_mode: 1},其中cube_nbuffer_setting/vec_nbuffer_setting控制 cube/vector 流水 buffer 数量auto_mix_partition开启自动混合切分device_sched_mode控制设备侧调度模式属于 PyPTO 算子调优的通用手段。5.3 Kernel 内部实现kernel 的计算组织与 README 描述的流程一一对应Grouped matmulexpert 并行token_num m // num_experts通过pypto.loop(config.num_experts, parallelTrue)按 expert 并行执行for expert_idx in pypto.loop(config.num_experts, parallelTrue): start expert_idx * token_num end (expert_idx 1) * token_num pypto.experimental.set_operation_options(combine_axisTrue) x_tile x1[start:end, :] pertoken_scale_tile pertoken_scale[start:end, :, :] weight_tile x2[expert_idx, :, :] weight_tile.set_cache_policy(pypto.CachePolicy.NONE_CACHEABLE, True) mm_result pypto.scaled_mm( x_tile, weight_tile, pypto.DT_FP32, pertoken_scale_tile, scale[:, :, :], a_transFalse, scale_a_transFalse, b_transconfig.transpose_x2, scale_b_transconfig.transpose_x2, ) gmm_out[start:end, :] mm_result要点x1按 expert token 范围连续切片x2按 expert 维度读取单个 expert 权重并标记为不可缓存以节省 L2 资源scaled_mm在 cube 上输出 FP32 中间结果b_trans与scale_b_trans随transpose_x2联动。Logit 加权与路由回写回写路径按route_tile 512分块串行执行parallelFalse每块做 unsqueeze、广播乘、index_add_三步if config.has_logit: for tile_idx in pypto.loop(route_tile_num, parallelFalse): result_tile gmm_out[start:end, :] logit_2d pypto.unsqueeze(logit[start:end], -1) result_tile pypto.mul(result_tile, logit_2d) pypto.index_add_(out, 0, row_index[start:end], result_tile)route_tile_num m // 512之外的尾部route_tail m % 512单独处理has_logitFalse时跳过乘 logit 分支直接index_add_。Shared expert 叠加kernel 内完成 cast 到 FP32、按shared_input_weight缩放、再index_add_到outif config.has_shared_input: shared_fp32 pypto.cast(shared_input[:, :], pypto.DT_FP32) shared_scaled pypto.mul(shared_fp32, config.shared_input_weight) pypto.index_add_(out, 0, shared_row_index, shared_scaled)其中shared_row_index在 host 侧由torch.arange(shared_input.shape[0]) shared_input_offset生成因此 README 中shared input 在 host 侧加到 out 初值、避免 kernel 内额外分支的设计在代码中的落地方式是host 侧预生成行索引、kernel 内统一走index_add_。5.4 Host 侧封装 gen_pyptogen_pypto(inputs)负责数据搬运与 kernel 启动把各输入搬到 NPUx1.npu()等、group_list转 CPU list 传入、row_index转 int32、预分配 FP32 的gmm_out中间缓冲区out深拷贝后作为累加初值最后调用gmm_finalize_routing_kernel并返回 FP32 的out。5.5 实现特点小结Cube Vector 融合scaled_mm使用 cube 计算logit 加权和 scatter-addindex_add_使用 vector 路径Expert 并行通过pypto.loop(config.num_experts, parallelTrue)按 expert 并行执行分块配置显式化使用set_cube_tile_shapes和set_vec_tile_shapes显式控制 cube/vector tile且 cube tile 由FinalizeRoutingConfig按M//E自动推导路由回写分块串行route 阶段按 512 行分块串行index_add_避免并行回写同一行带来的竞争。内存访问模式上x1按 expert token 范围连续切片x2按 expert 维度读取单个权重out通过row_index执行非连续 scatter-addshared input 在 host 侧预处理行索引。六、精度验证6.1 容差设置相对容差RTOL0.001绝对容差ATOL0.001容差在 test_gmm_finalize_routing.py 中定义为模块级常量RTOL 1e-3、ATOL 1e-3最终通过numpy.testing.assert_allclose(golden, result, rtolRTOL, atolATOL)校验。6.2 测试用例测试名称batchMKNEDType说明case11287686144409632FP8 E4M3大 K、32 expertscase22567688192409632FP8 E4M3更大 K、batch256case364128512040968FP8 E4M3小 M、8 expertscase4642567168409616FP8 E5M2FP8 E5M2 路径这四个 case 覆盖了大 K 大 Ecase1/case2、小 M 小 Ecase3以及 E5M2 输入格式case4三类典型组合与动态轴覆盖范围 {batch: 64/128/256, M: 128/256/768, K: 5120/6144/7168/8192, N: 4096, E: 8/16/32} 一一对应。说明单测文件中实际注册的TEST_CONFIGS为两个配置batch4, topk8, k7168, n4096, e8与batch128, topk8, k7168, n4096, e8均带pytest.mark.soc(950)。README 表格中的 case1~case4 描述了更广的覆盖目标实际以TEST_CONFIGS注册的用例为准。6.3 验证方法Golden 实现gmm_finalize_routing_golden.py 中的gen_golden。其_compute_mxfp8_matmul_golden以纯 PyTorch 方式复现 MXFP8 反量化按 K 维 32 元素对 E8M0FNU scale 做repeat_interleave展开成逐元素 scale再执行 FP32 matmulgen_golden遍历每个 expert、完成 logit 加权、index_add_回写与 shared_input 叠加。PyPTO 实现gmm_finalize_routing_impl.py 中gen_pypto调用gmm_finalize_routing_kernel。对比工具numpy.testing.assert_allcloseRTOL/ATOL 均为 1e-3。6.4 运行方式单测文件既支持 pytest也支持直接以脚本运行# 方式一pytest 运行全部用例 pytest tests/ops/experimental/matmul/grouped_matmul_finalize_routing/test_gmm_finalize_routing.py # 方式二直接运行不带参数执行全部用例 python tests/ops/experimental/matmul/grouped_matmul_finalize_routing/test_gmm_finalize_routing.py # 方式三直接运行按 1-based 序号执行单个用例 python tests/ops/experimental/matmul/grouped_matmul_finalize_routing/test_gmm_finalize_routing.py 1测试通过时打印description PASSED。运行前需保证环境已安装pypto、torch、torch_npu并处于对应 NPU 环境测试带有soc(950)标记。七、小结grouped_matmul_finalize_routing 展示了 PyPTO 编写 MoE 后处理融合算子的典型范式FinalizeRoutingConfig承载 Shape 与行为开关并自动推导 cube tilepypto.scaled_mm完成 MXFP8 分组矩阵乘pypto.loop(..., parallelTrue)实现 expert 并行pypto.unsqueeze/pypto.mul/pypto.index_add_组合完成 logit 加权与 scatter-add 回写shared expert 叠加通过 host 预生成行索引、kernel 内统一index_add_落地。配合纯 PyTorch 的 golden 实现与assert_allclose容差校验形成文档语义 → kernel 实现 → golden 对照的完整闭环可作为后续同类融合算子的参考模板。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻