pypto-gym 中的 SK-14 通用多矩阵乘骨架:PyPTO 多 MatMul 算子设计的兜底范式
pypto-gym 中的 SK-14 通用多矩阵乘骨架PyPTO 多 MatMul 算子设计的兜底范式【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gymSK-14General Multi-MatMul是 PyPTO 算子设计方法论中的“兜底骨架”当一个算子以多个矩阵乘为核心计算却对不上 Flash Attention、FFN/SwiGLU、MoE 等专用骨架时就用它来组织计算图。本文以 SK-14-general-matmul.md 的设计文档为主体结合 pypto-gym 仓库中 GLMGate、GMMFinalizeRouting、GMM-MXFP8、QuantMatMulReduceSum 等真实算子源码讲透这个骨架的适用判断、阶段编排、编码约束与性能调优配置读完你可以独立完成一个多 matmul 融合算子的 PyPTO kernel 设计与实现。一、SK-14 在骨架体系中的位置pypto-gym 的算子设计技能包把整体计算骨架编号为 SK-01 ~ SK-16覆盖 attention、projection、FFN、MoE、vector、recurrent、cache 等典型形态索引见 skeletons/index.md。其中 SK-14 的条目为ID名称tagsflow_patternSK-14General Multi-MatMulmatmulC, V这里的C 表示 Cube矩阵乘单元V 表示 Vector向量单元。SK-14 的定位是“通用兜底”适用场景以多个 matmul 为核心计算的算子且不匹配 SK-01~SK-07 等特定骨架。涵盖分组矩阵乘、量化矩阵乘、多路并行线性投影、多专家批量 matmul、matmul 间夹少量向量操作dequant/quant/reduce/scatter等。CV 排布多个 C 阶段显式展开C 之间可夹少量 V 操作模式为C...[V]...C...[V]...C。适用条件选型判据原文档给出的判断标准可以直接作为设计 checklisthas_matmul True且matmul_count 1不匹配 SK-01(FA)、SK-02(SinglePass)、SK-03(NormLinear)、SK-04(Prolog)、SK-06(FFN)、SK-07(MOE) 等特定骨架典型特征多个分组/并行/串行 matmulC 之间夹少量 V 操作非完整 CV 融合管线与 SK-15General CV Fusion的区别SK-14 的 C 阶段占主导V 阶段仅做量化/反量化/激活等轻量后处理若向量计算本身很重例如完整的归一化激活量化管线应走 SK-15。二、编码约束JIT 需要编译期可见的完整计算图SK-14 文档中有一条硬性编码约束是多 matmul 算子开发中最容易踩坑的地方禁止使用for mm_idx in range(N)optional_v_stage()循环分发 matmul每个 C/V 阶段必须显式写出JIT 编译器需要在编译期确定完整计算图。也就是说pypto.frontend.jit装饰的 kernel 里每一个 matmul 阶段、每一个向量阶段都要在源码中静态可见允许被 JIT 静态展开的循环但不允许“运行时才决定跑哪几个 matmul”的动态分发。这一点在仓库源码中可以得到印证gmm_mxfp8_impl.py 的 docstring 明确解释了为什么分组循环要用编译期可静态求值的range()config.group_list累加 begin/end 来展开而不是运行时动态注解——显式的pypto.Tensor([...], ...)注解会使b.shape[0]变为 SymbolicScalar导致range()与config.group_list[i]编译失败。三、骨架结构完整代码模板以下是原文档给出的标准骨架每个可选阶段都必须独立设置 tile shapespypto.frontend.jit( pass_options{ cube_l1_reuse_setting: {-1: 8}, cube_nbuffer_setting: {-1: 4}, vec_nbuffer_setting: {-2: 1, -1: 8}, }, runtime_options{stitch_function_max_num: 128}, ) def general_multi_matmul_kernel(A, B_list, output, config): pypto.experimental.set_operation_options(combine_axisTrue) # ── ○ 可选 静态预处理 ───────────────────────── pypto.set_vec_tile_shapes(v_static_tile) # 预处理阶段: 独立设置 ... # 动态轴提取 / 参数 cast / 常量计算 # ── ○ 可选 Shape 前置变换 ───────────────────── pypto.set_vec_tile_shapes(v_shape_tile) # Shape 变换阶段: 独立设置 ... # 3D→2D reshape / transpose / 合轴拆轴 / 静态 shape 推导 # ── ○ 可选 外层 Loop ───────────────────────── # 变体 A: 无 loop全量串行 matmul # 变体 B: parallel loop → for g_idx in pypto.loop(num_groups, parallelTrue): # 变体 C/D: 串行 loop → for g_idx in pypto.loop(tile_count): ... # 分组切片: a_tile group_slice(A, g_idx) # ── ○ 可选 C_1 ─────────────────────────────── pypto.set_cube_tile_shapes(c1_tile) # 每个 MatMul 必须独立设置 ... # MatMul (如 Gate Projection / QKV Proj / QuantMatMul) # ── ○ 可选 V_inter ──────────────────────────── pypto.set_vec_tile_shapes(v1_tile) ... # dequant / quant / split / activation保持轻量重 V 操作走 SK-15 # ── ○ 可选 C_2 ─────────────────────────────── pypto.set_cube_tile_shapes(c2_tile) # 必须重设M/N/K 通常不同 ... # MatMul (如 Up Projection / Second Proj) # ── ○ 可选 V_inter ──────────────────────────── pypto.set_vec_tile_shapes(v2_tile) ... # SwiGLU / GELU / concat / cast # ── ○ 可选 C_N (可重复 1~N 次) ──────────────── pypto.set_cube_tile_shapes(cN_tile) # 每个 MatMul 前必设 ... # MatMul (如 Down Projection / Output Proj) # ── ○ 可选 V_post后聚合──────────────────── pypto.set_vec_tile_shapes(v_post_tile) ... # reduce / cast / scatter / index_put_ 写回 # ── ○ 可选 结果存储 ─────────────────────────── ... # group_store / index_put_(accumulateTrue) / assemble 写回注意模板首行的pypto.experimental.set_operation_options(combine_axisTrue)文档将其列为“必配”放在 jit 函数体首行用于开启算子选项的合轴优化。四、变体速查四种 Loop 结构原文档把 SK-14 实例归纳为四种变体并给出了对应仓库实例变体Loop 结构MatMul 模式典型场景实例A: 无 Loop无显式 loop2 串行 matmul全量多路投影管线QuantMatMulReduceSumB: 并行 Looploop(N, parallelTrue)每组 1~2 个 matmul多专家分组并行GMMFinalizeRouting, QuantGroupedMMC: 串行 Looploop(tiles)每组 1 个 matmul逐 tile 批量投影GLMGateD: 串行多 MatMulloop(tiles)每组 2~3 个 matmulgateup 并行→downPanguFusedLayer(FFN 部分), FusedSwiGLU五、四种典型编排模式SK-14 文档总结了多 matmul 计算的四种数据流拓扑这也是设计 kernel 时最实用的组织框架# 模式 1: 串行管线前一 matmul 输出 → 后一 matmul 输入 C1(A, W1) → [V: dequant] → C2(mm1_out, W2) → [V: dequant] → C3(mm2_out, W3) 实例: 多级量化投影 (MLA: q_a_proj → norm → q_b_proj) # 模式 2: 并行多路同一输入 → 多个独立 matmul → 合并 C1(A, W_gate) ─┐ ├→ [V: SwiGLU/concat] → C3(merged, W_down) C2(A, W_up) ──┘ 实例: FFN (gate_proj up_proj → SwiGLU → down_proj) # 模式 3: 分组并行不同输入分片 → 各自 matmul → scatter/gather for g in parallel: C_g(A[g], W[g]) → [V: scale] → scatter(output) 实例: GMM (per-expert matmul logit scaling) # 模式 4: 循环内多 CKV 分片场景每组多步 matmul for group: C1(Q, K[group]) → [V: softmax] → C2(P, V[group]) 实例: 通用稀疏 attention非标准 FA 场景这四种模式在仓库源码中都有对应实现下面逐一结合代码解析。变体 B 实战GMMFinalizeRouting分组并行 后聚合 Vgrouped_matmul_finalize_routing 是模式 3分组并行的完整实现MoE 场景下按 expert 分组的 MXFP8 矩阵乘 logit 加权 按 row_index 累积写回 shared_input 叠加。其 kernel 结构与 SK-14 骨架逐段对应JIT 配置L121-L129pass_options配cube_nbuffer_setting: {-1: 1}、vec_nbuffer_setting: {-2: 1, -1: 1}、auto_mix_partition: 1runtime_options配stitch_function_max_num: 128与device_sched_mode: 1并行调度——与文档“多 MatMul 并行调度推荐 device_sched_mode1”的建议一致。并行分组循环L165-L188for expert_idx in pypto.loop(config.num_experts, parallelTrue)每组内先按start:end切片输入与 per-token scale再调用pypto.scaled_mm(x_tile, weight_tile, pypto.DT_FP32, pertoken_scale_tile, scale, ...)完成 MXFP8 量化矩阵乘结果直接gmm_out[start:end, :] mm_result写回。循环内还调用了pypto.experimental.set_operation_options(combine_axisTrue)。post-loop V 聚合L190-L225循环结束后一次性重设 vec tile shapes按 512 行的 tile 串行循环做 logit 广播乘法并通过pypto.index_add_(out, 0, row_index[start:end], result_tile)把多组结果累积写回同一输出尾部route_tail单独处理即“tail block”模式见 AT-20最后把shared_inputcast 成 FP32、乘权重再index_add_叠加。这正是文档强调的两条关键编码特征index_put_/index_add_ 累积写与post-loop V 一次性聚合避免循环内重复 V 启动。其 host 侧封装gen_pyptoL230-L265展示了调用约定输入张量.npu()上卡、group_list转 CPU list 传入 kernel编译期静态分组、预分配 FP32 的gmm_out承接中间结果。配置数据类FinalizeRoutingConfigL27-L90里按per_expert_m分档在__post_init__中自动推导 M/K/N tile shapes是“tile 配置与形状联动”的实际写法。变体 B 的另一形态GMM-MXFP8 的编译期静态分组gmm_mxfp8_impl.py 实现了“输入按 group_list 切行、每组乘不同权重矩阵再拼回一个输出”的分组 GEMMmatmul/README.md 中给出了计算公式Output Concat(A_1·W_1, ..., A_g·W_g)与 MXFP8 布局数据 E4M3FN、缩放因子 E8M0FNU、每 64 个元素共享一个缩放因子。kernel 主体L130-L155的写法round_num b.shape[0] begin 0 end 0 ... for i in range(round_num): begin end end end config.group_list[i] x a[begin:end, :] weight b[i] scaled_x scaled_a[begin:end, :, :] scaled_weight scaled_b[i] out[begin:end, :] pypto.scaled_mm( x, weight, pypto.DT_FP32, scaled_x, scaled_weight )两个值得注意的点这里的for i in range(round_num)不是运行时分发而是依赖round_num b.shape[0]在 JIT 编译期可静态求值、由编译器展开的 Python 循环——正对应 SK-14 编码约束“每个 C 阶段必须显式写出”文件 docstringL16-L35还记录了该 kernel 使用空注解pypto.Tensor()而非显式 shape 注解的编译期原因。scaled_mm(A, B, FP32, scale_a, scale_b)就是文档“关键编码特征”表中强调的 MXFP8 统一接口量化信息随矩阵乘一起进入 Cube 计算不要拆出独立 quant/dequant 阶段。该 kernel 的 JIT 配置L102-L111为cube_nbuffer_setting: {-1: 4}、vec_nbuffer_setting: {-2: 1, -1: 4}、device_sched_mode: 3说明文档给出的 buffer/sched 配置应视为“起点经验值”按实际编译规模与资源占用调整。变体 C 实战GLMGate 的逐 tile 串行 matmulglm_gate_impl.py 是 GLM-4.5 MoE 的 gate 投影算子hidden_states 从 d_model 5120 维投影到 d_router 160 维对应文档变体 C串行 Loop、每组 1 个 matmul、逐 tile 批量投影。其 kernel 核心L77-L96展示了动态 batch 下 tile 串行投影的标准套路bs hidden_states.shape[0] h_num hidden_states.shape[1] view_shape (32, h_num) bs_loop (bs view_shape[0] - 1) // view_shape[0] for bs_idx in pypto.loop(bs_loop, nameLOOP_MOE_MM_L0, idx_namebs_idx): tile_hidden_states pypto.view(hidden_states, view_shape, [bs_idx * view_shape[0], 0], valid_shape[(bs - bs_idx * view_shape[0]).min(view_shape[0]), h_num]) pypto.set_cube_tile_shapes([32, 32], [512, 1024], [16, 16]) res pypto.matmul(tile_hidden_states, mm_weight, tile_hidden_states.dtype, b_transTrue) router_logits_out[bs_idx * view_shape[0]:, 0:] res这里体现了 SK-14 两条关键编码特征的落地K 轴/M 轴分组切分按 32 行一组pypto.view(..., valid_shape...)切片最后一组的valid_shape用.min()处理不满 tile 的尾块即文档“K 轴分组切分view valid_shape 按组切片见 AT-20”每 matmul 独立 TileShapeset_cube_tile_shapes写在循环内、每次迭代显式设置避免复用前一个阶段的 tile 配置。host 侧的gate函数L101-L133带allow_in_graph装饰器并做了 FakeTensor 旁路返回预分配输出说明该算子面向 torch.compile 图集成输入校验check_args要求权重与激活均为 FP32、ND 格式这是该算子当前实现的适用前提。变体 A 实战QuantMatMulReduceSummatmul 轻量 V_postquant_matmul_reduce_sum_impl.py 对应文档变体 A无显式 loop 的批量 matmul 后聚合 V对[batch, m, k] × [batch, k, n]做 INT8 批量矩阵乘反量化后沿 batch 维 reduce sum。kernelL75-L110的 C→V 编排非常典型pypto.set_cube_tile_shapes(config.m_tile_shape, config.k_tile_shape, config.n_tile_shape) pypto.set_vec_tile_shapes(*config.vec_tile_shapes) if config.x2_format_nz: pypto.set_matrix_size([m, k, n]) # C 阶段INT8 MatMul输出 INT32 matmul_result pypto.matmul(x1, x2, pypto.DT_INT32) matmul_result_fp32 pypto.cast(matmul_result, pypto.DT_FP32) # V 阶段scale 广播 反量化乘法 reduce sum x2_scale_fp32 pypto.cast(x2_scale, pypto.DT_FP32) x2_scale_2d pypto.unsqueeze(x2_scale_fp32, 0) x2_scale_broadcast pypto.expand_clone(x2_scale_2d, [m, n]) x1_scale_2d pypto.unsqueeze(x1_scale, 2) x1_scale_broadcast pypto.expand_clone(x1_scale_2d, [batch, m, n]) scale_mul pypto.mul(x1_scale_broadcast, x2_scale_broadcast) scaled pypto.mul(matmul_result_fp32, scale_mul) out_fp32 pypto.sum(scaled, 0) out_bf16 pypto.cast(out_fp32, pypto.DT_BF16) out.move(out_bf16)该例同时覆盖了文档“关键编码特征”表中的两项NZ 格式权重x2声明为pypto.TileOpFormat.TILEOP_NZL73且 NZ 场景下调用pypto.set_matrix_size([m, k, n])声明逻辑矩阵尺寸——对应文档“权重可能使用 NZ(fractal) 格式需在签名中声明”。V_post 保持轻量matmul 之后的向量阶段只做 cast/expand/mul/sum没有任何重组或重依赖正好落在 SK-14 允许的 V 操作集合内。matmul 目录 README 给出了该算子的完整参数表x1:[batch, M, K]INT8x2:[batch, K, N]INT8ND/NZx1_scale:[batch, M]FP32x2_scale:[N]BF16输出[M, N]BF16及五步计算公式可作为变体 A 算子规格书的参考模板。六、阶段积木速查与关键编码特征阶段积木积木类型可选操作可重复依赖C_NCMatMul (BF16 / INT8 / MXFP8 via scaled_mm)是 (1~N)—V_interVdequant / quant / split / activation是相邻 C 阶段V_postVreduce / cast / scatter / index_put_否最后一个 C 阶段V_inter 必须保持轻量仅 quant/dequant/split/activation重 V 操作应走 SK-15。关键编码特征特征说明仓库印证每 matmul 独立 TileShape不同 matmul 的 M/N/K 维度通常不同需各自set_cube_tile_shapes复用会触发表达式爆炸glm_gate_impl.py#L92 在每次迭代显式设置parallelTrue分组独立时启用并行循环编译器自动并行化gmm_finalize_routing_impl.py#L165K 轴分组沿 K/M 轴按组切分 A/B每组独立 matmulgmm_mxfp8 按config.group_list切行scaled_mmMXFP8 量化 matmul 的统一接口scaled_mm(A, B, FP32, scale_a, scale_b)gmm_mxfp8 / gmm_finalize_routingindex_put_ 累积多组写入同一输出 tensor 时使用accumulateTrue或index_add_gmm_finalize_routing 的index_add_写回C 间 V 保持轻量两个 matmul 之间的 V 操作仅做 dequant/quant/activation重 V 走 SK-15各实例均遵守post-loop Vloop 后一次性做 batch add / reduce / castgmm_finalize_routing 的 logit 加权与 shared_input 叠加NZ 格式权重权重可能使用 NZ(fractal) 格式需在签名中声明quant_matmul_reduce_sum 的TILEOP_NZ七、开箱性能优化配置文档推荐表全量继承推断来源GLMGate/GMMRouting/QuantGroupedMM 等多 MatMul 算子共性MXFP8 量化通过scaled_mm接口。维度推荐配置取值经验作用pass_options.cube_l1_reuse_setting必配多键分轴{-1: 2, 0: 4, 1: 1}每个 matmul 的 L1 复用策略不同需精细pass_options.cube_nbuffer_setting必配{-1: 2}起步分组 MM 升{1: 2}多 cube 阶段双缓冲pass_options.vec_nbuffer_setting推荐{-1: 2}即可V 阶段轻量C 间 V 仅做量化等轻操作runtime_options.stitch_function_max_num必配128限制 stitch 函数规模runtime_options.device_sched_mode推荐1并行多 MatMul 并行调度每个 MatMul 独立set_cube_tile_shapes强制每个 matmul 前必设M/N/K 维度不同复用前一个的 tile 会触发表达式爆炸parallelTrue推荐分组循环专家/group编译器自动并行化K 轴分组切分推荐view valid_shape按组切片见 AT-20scaled_mm接口必配量化MXFP8 用统一接口不要拆出独立 quant/dequantindex_put_(target, idx, v, accumulateTrue)推荐多组写同一输出比拆 assemble 高效C 间 V 操作保持轻量仅 quant/dequant/split/activation重 V 操作应分骨架做SK-15post-loop V 聚合推荐循环后一次性add/reduce/cast避免循环内重复 V 启动NZ 权重格式推荐签名声明formatNZ提升 cube 加载效率combine_axisTrue必配jit 首行开启合轴算子选项对照仓库实现可以看到这张表是“起点”而非“死值”gmm_finalize_routing 用cube_nbuffer_setting: {-1: 1}、device_sched_mode: 1gmm_mxfp8 用cube_nbuffer_setting: {-1: 4}、vec_nbuffer_setting: {-2: 1, -1: 4}、device_sched_mode: 3而 glm_gate 的runtime_options为空。可以推断这些差异来自各算子的分组规模、编译产物体积与目标芯片调度策略。文档给出的通用调优建议是性能建议分别检查各矩阵乘的 TileShape。分组之间没有数据依赖时可以评估parallelTrue比较编译规模、资源占用和运行耗时。八、验证与继续深入SK-14 各实例都配有 Golden 参考实现与单测验证路径如下GMMFinalizeRouting / GMM-MXFP8 / QuantMatMulReduceSumtests/ops/experimental/matmul/grouped_matmul_finalize_routing/、tests/ops/experimental/matmul/gmm_mxfp8/、tests/ops/experimental/matmul/quant_matmul_reduce_sum/各目录内的 golden 与 test 入口见 matmul/README.md 的算子说明与参数表GLMGatetests/ops/glm_v4_5/test_glm_gate.py配套实现与集成说明见 glm_v4_5 目录。设计方法论文档方面可配合阅读骨架总索引 skeletons/index.mdSK-14 与 SK-01~SK-16 的完整选型表尾块处理原子 AT-20-tail-block.mdview valid_shape切片的完整规则重向量融合骨架 SK-15-general-cv-fusion.md当 C 间 V 阶段变重时的分流出口。总结SK-14 的价值在于给“多 matmul 但不匹配任何专用骨架”的算子提供了一个可执行的设计协议先用适用条件判断是否落在 SK-14再按“四种变体 × 四种编排模式”确定 Loop 结构与数据流然后逐阶段显式写出 C/V 积木并各自设置 tile shapes最后按开箱配置表调 buffer/调度参数并比较编译规模与运行耗时。仓库中的 GMMFinalizeRouting、GMM-MXFP8、GLMGate、QuantMatMulReduceSum 四个实现正好覆盖了并行分组、编译期静态分组、动态 batch 尾块、批量后聚合这四类典型形态是学习这一骨架最直接的一组参考样本。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考