基于 PyPTO 的 MoE 融合算子实践:grouped_matmul_finalize_routing 的 MXFP8 实现与精度验证
【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym
grouped_matmul_finalize_routing 是 CANN pypto-gym 仓库中面向 MoE(Mixture 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_offset+batch] += shared_input * shared_input_weight展开形式为:
out[row_index[t], n] += logit[t] * Σ(k=0..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 计算流程
整个算子按以下四个阶段执行:
阶段 1:Grouped 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]
阶段 2:Logit 加权。当has_logit=True时,对每个 token 的 matmul 结果乘以对应logit:
logit_i: [M_i]→ unsqueeze →[M_i, 1]- 广播乘法:
[M_i, N] × [M_i, 1] → [M_i, N]
阶段 3:Finalize Routing 回写。根据row_index将 expert 输出累加到最终输出:
row_index_i: [M_i]out[row_index_i] += weighted_i
阶段 4:Shared Expert 叠加。当has_shared_input=True时,将 shared expert 输出按权重叠加到out:
shared_input: [batch, N]out[offset:offset+batch] += shared_input * shared_input_weight
三、输入输出规格
3.1 输入张量
| 名称 | Shape | DType | 说明 |
|---|---|---|---|
x1 | [M, K] | FP8 E4M3/E5M2 | 路由后 token 输入 |
x2 | [E, K, N]或[E, N, K] | FP8 E4M3/E5M2 | expert 权重,布局由transpose_x2决定 |
scale | [ceil(K/64), N, 2]或[N, ceil(K/64), 2] | E8M0FNU | 权重 scale |
pertoken_scale | [M, ceil(K/64), 2] | E8M0FNU | token scale |
group_list | [E] | int64 | expert 分组信息 |
shared_input | [batch, N] | bfloat16 | shared expert 输出 |
logit | [M] | float32 | token 对应 expert 权重 |
row_index | [M] | int64 | 输出行索引 |
out | [batch, N] | float32 | 输出初值 |
3.2 输出张量
| 名称 | Shape | DType | 说明 |
|---|---|---|---|
output | [batch, N] | float32 | finalize 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_e5m2;transpose_x2=True时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) // 64;row_index = torch.arange(config.m) % config.batch,保证索引落在[0, batch);shared_input使用 bfloat16,out为 FP32 全零初值。
四、Shape 范围与约束
4.1 动态轴(当前覆盖范围)
| 轴 | 当前覆盖范围 | 说明 |
|---|---|---|
| batch | {64, 128, 256} | 输出行数 |
| M | {128, 256, 768} | 路由后 token 数 |
| K | {5120, 6144, 7168, 8192} | matmul K 维 |
| N | 4096 | 输出列数 |
| E | {8, 16, 32} | expert 数量 |
4.2 约束条件
- transpose_x1 仅支持 False:当前目标路径不支持转置
x1。这一点在 golden 中也被强制校验——gen_golden中若cfg.transpose_x1为 True 会直接raise ValueError("aclnnGroupedMatmulFinalizeRoutingV3 only supports transposeX1=False.")。 - 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,当前测试取值为 1):
group_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 数;m:token 总数,由batch * topk自动计算(不可手动指定);k/n:matmul 的 K 维与 N 维(输出列数);num_experts:expert 数量;in_dtype:输入 FP8 数据类型,默认pypto.DT_FP8E4M3;transpose_x1/transpose_x2:是否转置输入(transpose_x1当前仅支持 False);group_list_type:0=前缀和,1=每组计数;shared_input_weight(默认 1.0)与shared_input_offset(默认 0):shared_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 tile:
per_expert_m <= 64:m_tile_shape=[per_expert_m, per_expert_m],k_tile_shape=[256, 512],n_tile_shape=[256, 512];per_expert_m <= 1024:m_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 matmul(expert 并行):token_num = m // num_experts,通过pypto.loop(config.num_experts, parallel=True)按 expert 并行执行:
for expert_idx in pypto.loop(config.num_experts, parallel=True): start = expert_idx * token_num end = (expert_idx + 1) * token_num pypto.experimental.set_operation_options(combine_axis=True) 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_trans=False, scale_a_trans=False, b_trans=config.transpose_x2, scale_b_trans=config.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分块串行执行(parallel=False),每块做 unsqueeze、广播乘、index_add_三步:
if config.has_logit: for tile_idx in pypto.loop(route_tile_num, parallel=False): 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_logit=False时跳过乘 logit 分支,直接index_add_。
Shared expert 叠加:kernel 内完成 cast 到 FP32、按shared_input_weight缩放、再index_add_到out:
if 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_pypto
gen_pypto(inputs)负责数据搬运与 kernel 启动:把各输入搬到 NPU(x1.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-add(index_add_)使用 vector 路径; - Expert 并行:通过
pypto.loop(config.num_experts, parallel=True)按 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-add,shared input 在 host 侧预处理行索引。
六、精度验证
6.1 容差设置
- 相对容差(RTOL):0.001
- 绝对容差(ATOL):0.001
容差在 test_gmm_finalize_routing.py 中定义为模块级常量RTOL = 1e-3、ATOL = 1e-3,最终通过numpy.testing.assert_allclose(golden, result, rtol=RTOL, atol=ATOL)校验。
6.2 测试用例
| 测试名称 | batch | M | K | N | E | DType | 说明 |
|---|---|---|---|---|---|---|---|
| case1 | 128 | 768 | 6144 | 4096 | 32 | FP8 E4M3 | 大 K、32 experts |
| case2 | 256 | 768 | 8192 | 4096 | 32 | FP8 E4M3 | 更大 K、batch=256 |
| case3 | 64 | 128 | 5120 | 4096 | 8 | FP8 E4M3 | 小 M、8 experts |
| case4 | 64 | 256 | 7168 | 4096 | 16 | FP8 E5M2 | FP8 E5M2 路径 |
这四个 case 覆盖了"大 K 大 E"(case1/case2)、"小 M 小 E"(case3)以及 E5M2 输入格式(case4)三类典型组合,与动态轴覆盖范围 {batch: 64/128/256, M: 128/256/768, K: 5120/6144/7168/8192, N: 4096, E: 8/16/32} 一一对应。
说明:单测文件中实际注册的
TEST_CONFIGS为两个配置(batch=4, topk=8, k=7168, n=4096, e=8与batch=128, topk=8, k=7168, n=4096, e=8,均带@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 matmul;gen_golden遍历每个 expert、完成 logit 加权、index_add_回写与 shared_input 叠加。 - PyPTO 实现:gmm_finalize_routing_impl.py 中
gen_pypto调用gmm_finalize_routing_kernel。 - 对比工具:
numpy.testing.assert_allclose(RTOL/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 tile;pypto.scaled_mm完成 MXFP8 分组矩阵乘;pypto.loop(..., parallel=True)实现 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),仅供参考