MiniMax MoE / MSA 融合算子实战指南:基于 PyPTO 的 M2.7 / M3 文本骨干网络算子实现与验证
2026/9/19 18:34:31 网站建设 项目流程

MiniMax MoE / MSA 融合算子实战指南:基于 PyPTO 的 M2.7 / M3 文本骨干网络算子实现与验证

【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym

导读

本文是 CANN pypto-gym 仓库中 MiniMax MoE / MSA Fused Operators (M2.7 / M3) 文档 的完整技术解读。该目录承载了基于 PyPTO 编程框架在 Ascend NPU 上实现的两类 MiniMax 文本骨干网络融合算子:M2.7 与 M3 共用的 MoE Grouped GEMM,以及 M3 专有的 MSA(MiniMax Sparse Attention)三件套——lightning-indexer、block-sparse decode attention 与 Main Branch GQA-batched flash attention。读完本文,你将掌握四个算子的算法结构、循环与 tile 设计、权重布局、Kernel 签名、环境变量旋钮,以及对应的精度测试与运行命令,并能在 tests/ops/minimax_m27/ 与 tests/ops/minimax_m3/ 中直接复现验证。

四个算子的实现源码位于 src/pypto_gym/ops/pypto_tensor/minimax/ 目录,入口汇总如下:

算子入口函数实现文件适用模型说明
minimax_moe_grouped_gemmminimax_moe_grouped_gemm()minimax_grouped_gemm_impl.pyM2.7 / M3All-experts single grouped GEMM,通过activation参数区分变体
minimax_m3_msa_indexerminimax_m3_msa_indexer()minimax_m3_msa_indexer_impl.pyM3MSA lightning indexer:block 选择
minimax_m3_msa_sparse_decodeminimax_m3_msa_sparse_decode()minimax_m3_msa_sparse_attention_impl.pyM3MSA block-sparse decode attention
msa_main_branchmsa_main_branch()msa_main_branch_impl.pyM3MSA Main Branch GQA-batched flash attention

产品支持情况与依赖

该算子包面向 Ascend NPU 平台,官方支持情况如下:

  • Ascend 910B:支持(Grouped GEMM 的 tile 默认值即针对 910B 的 CUBE/VECTOR 分核架构与 192 KB UB 调优)
  • Ascend 950 / 950PR:支持

运行环境依赖(与 modeling/transformers/minimax/README.md 中记录的 M2.7 / M3 迁移环境一致):

  • Python 3.x
  • PyTorch + torch_npu(迁移环境记录为 torch_npu 2.10)
  • PyPTO(pypto包,迁移环境记录为 0.2.1,配套 pto-isa v9.1.0、CANN 9.1.0)
  • NumPy
  • pytest(可选,用于 skip 标记与参数化)

算子一:minimax_moe_grouped_gemm(M2.7 / M3 共用)

算法概述

MiniMax MoE expert FFN 的融合 grouped GEMM:一次 kernel 调用完成所有 expert 的mm1 → activation → mm2pypto.loop遍历 expert,pypto.loop_unroll遍历 token,配以 UB-fitting 的 vector tile。

两个 MiniMax 变体仅 expert activation 不同,通过activation参数选择:

变体activation激活公式tile 默认值 (VEC_TILE / CUBE_NBUF / VEC_NBUF / L1)
M2.7"silu"SiLU(gate) * up128 / 2 / 2 / 2
M3"swigluoai"(clamp(up) + 1) * (gate * sigmoid(alpha * gate))256 / 4 / 1 / 3

swigluoai 参数取自 M3 真实 config(text_config中的swiglu_alpha/swiglu_limit):alpha = 1.702limit = 7.0。在源码中对应 minimax_grouped_gemm_impl.py 的模块级常量_SWIGLU_ALPHA_SWIGLU_LIMIT_VARIANT_DEFAULTS字典,可用环境变量PYPTO_SWIGLU_ALPHA/PYPTO_SWIGLU_LIMIT覆盖。

两个激活的 kernel 内实现(同一处源码)分别由_swiglu_silu_swiglu_oai两个 Python 函数表达(见源码 L100-L123):

  • M2.7silupypto.viewgate_up切分为 gate 与 up 两半,pypto.exp计算exp(-gate),最终SiLU(gate) * up
  • M3swigluoai(GPT-OSS 风格 clamped GLU):gate = clamp(gate, max=limit)up = clamp(up, -limit, +limit)glu = gate * sigmoid(alpha * gate),输出(up + 1) * glu

从源码注释看,该 kernel 在结构上与llada2_moe的 grouped GEMM 共享同一套设计(minimax_grouped_gemm_impl.py)。

循环结构

EXPERT_LOOP — 遍历 expert,偏移从 expert_cumsum 动态获取 LOOP_TOKEN (unroll)— 每个 expert 的 token 按 unroll_list 分块 mm1: tile_x @ w13_e → gate_up [tile_batch, 2*I] FP32 activation (silu/swigluoai) → cast BF16 mm2: sw @ w2_e → down [tile_batch, H] FP32 → cast BF16 assemble → result

源码实现细节(minimax_grouped_gemm_impl.py):

  • 外层EXPERT_LOOP通过expert_cumsum[e_idx]/expert_cumsum[e_idx + 1]动态获取每个 expert 的 token 起止,得到该 expert 的 token 数n_e
  • 内层LOOP_TOKEN使用pypto.loop_unrollunroll_list分块,默认 unroll 列表由环境变量PYPTO_UNROLL控制(默认"1,2,4,8,16,32,64"),M3 的 pytest smoke 会将其收窄为"1,2,4"以控制规模;
  • w13_e/w2_e通过pypto.view从扁平权重中按 expert 偏移切出;
  • mm1 用pypto.set_cube_tile_shapes设置[tile_batch, tile_batch] / [_MM1_K, _MM1_K*2] / [_MM1_N, _MM1_N]pypto.matmul(tile_x, w13_e, pypto.DT_FP32)输出 FP32 的gate_up [tile_batch, 2*I]
  • activation 在 FP32 全程计算后pypto.cast(..., pypto.DT_BF16)回 BF16,vector tile 宽度被裁剪为min(tile_batch, vt)以适配 UB;
  • mm2 输出 FP32down [tile_batch, H],再 cast 到 BF16 后pypto.assemble写回 result。

mm1 / mm2 的 K、N 分块大小同样可配:PYPTO_MM1_K(默认 128)、PYPTO_MM1_N(默认 256)、PYPTO_MM2_K(默认 128)、PYPTO_MM2_N(默认 256)。

权重布局

F.linear 约定(输入): gate_up_proj [E, 2*I, H] — gate||up weights down_proj [E, H, I] — down weights convert_minimax_weights 转换后(direct matmul 格式): w13_flat [E*H, 2*I] BF16 — flattened gate||up w2_flat [E*I, H] BF16 — flattened down

convert_minimax_weights的实现(minimax_grouped_gemm_impl.py)非常简单:w13_flat = gate_up_proj.transpose(1, 2).reshape(E*H, 2*I)w2_flat = down_proj.transpose(1, 2).reshape(E*I, H),均.contiguous()保证内存连续。这与真实 MiniMax-M3 权重命名block_sparse_moe.experts.N.{w1,w2,w3}w1/w3拼为 gate||up,w2为 down)对应。

Kernel 签名

minimax_moe_grouped_gemm( sorted_tokens, # [N_total, H] BF16 — tokens pre-sorted by expert weights, # (w13_flat, w2_flat) BF16 — converted weights expert_cumsum, # [E+1] INT32 — cumulative token counts per expert result, # [N_total, H] BF16 — output buffer dims, # MoeDims(num_experts, hidden_size, intermediate_size, activation) )
  • MoeDims是定义在 minimax_grouped_gemm_impl.py 的NamedTupleactivation默认"silu"
  • 入口函数带@allow_in_graph装饰,FakeTensor输入直接返回(供 torch.compile 图模式使用);_check会对各张量的维数、dtype、形状做严格校验(如w13_flat必须是[E*H, 2I]的 BF16 2-D 张量);
  • 两个变体的 kernel 在 import 时通过_build_kernel各构建一次(jit 装饰器在 import 时廉价应用,TBE 编译推迟到首次真实调用),只有被当前模型实际调用的变体才会被编译(minimax_grouped_gemm_impl.py)。

Dtype 转换流程

阶段操作Dtype
输入sorted_tokens / w13_flat / w2_flatBF16
mm1tile_x @ w13_eBF16 → FP32 (out_dtype=FP32)
activationsilu / swigluoaiFP32 全程
中间 castactivation 输出 → BF16FP32 → BF16
mm2sw @ w2_eBF16 → FP32 (out_dtype=FP32)
输出 castdown → resultFP32 → BF16

即:BF16 输入输出、FP32 累加,中间 activation 结果回 BF16 以压缩 mm2 的 K 维带宽。值得注意的是,M3 的 golden 参考在 activation 输出处也做了 BF16 舍入(见 test_minimax_m3_grouped_gemm.py 注释),因为当激活值大到触发 clamp 时,FP32 中间结果与真实 dtype 流会产生分歧——这解释了为何 M3 用例的 atol 需要放宽到 0.5。

910B 适配要点:UB-fitting vector tile

源码 docstring(minimax_grouped_gemm_impl.py)明确说明:910B 是 CUBE/VECTOR 分核架构,UB 仅 192 KB,若 FP32 vector tile 取完整的 intermediate/hidden 宽度会溢出 UB 并导致OoOSchedulepass 失败;将宽度限制在[128,128]FP32 双缓冲(128 KB)即可正常调度并与参考对齐。这正是silu变体默认VEC_TILE=128的由来;swigluoai变体 H=6144/I=3072 更宽,因而采用VEC_TILE=256 / CUBE_NBUF=4 / VEC_NBUF=1 / L1=3的组合。

opt-in 开关:USE_PTO_GROUPED_GEMM

init.py 中声明USE_PTO_GROUPED_GEMM = False(opt-in,默认关闭),并重导出grouped_gemmconvert_minimax_weightsMoeDims。modeling 层读取该开关决定 MoE FFN 是否走融合 kernel(这是唯一开关;routing / gate 仍留在 host 端)。在 modeling/transformers/minimax/README.md 中,通过--use_pypto参数在推理与 benchmark 脚本中启用。M3 的 MSA decode kernel 由 modeling 层与测试直接从各自的*_impl模块导入,不在__init__.py重导出。

算子二:minimax_m3_msa_indexer

算法概述

M3 MiniMax Sparse Attention 的 lightning indexer(decode 步骤,单 query)。4 个 index-query head 对单条 (MQA) index key 全序列打分,每 128 token block 做 max-pool,再跨 4 head 取 max,最后 top-k 选出 16 个 block(含强制保留的 local block)。

PyPTO kernel 只负责计算量大的BF16 score matmul(O(ctx) 部分),将结果写入[NPAD, nb*BY]的 FP32 score buffer;而 block max-pool / head-max / top-k / local 强制这些廉价尾巴在 host 端 eager torch 的[nb]小向量上执行。

关键常量

常量说明
NIDX4sparse_num_index_heads
NPAD16cube M 轴 16 对齐,4 head 零填充到 16
D128sparse_index_dim
BY128sparse_block_size
TOPK16sparse_topk_blocks
LOCAL1sparse_local_block

cube 约束:matmul 的 M 轴必须 16 对齐,因此 4 个 index head 零填充到 16(NPAD),pad 行产生的 score 为 0,在 host 端通过[:NIDX]切片丢弃(源码注释见 minimax_m3_msa_indexer_impl.py)。

Kernel 签名

minimax_m3_msa_indexer( idx_q, # [NIDX, D] BF16 — index-query heads (post norm + RoPE) idx_k, # [nb*BY, D] BF16 — single index key over all keys (post norm + RoPE) nb, # int — number of 128-token key blocks ) # Returns: [1, min(TOPK, nb)] INT32 — selected block ids

实现要点:

  • kernel 按nb做 JIT 编译并按nb缓存(_kernel_cache),因为张量形状是静态的;kernel 内部对每个 block 静态展开for blk in range(nb),以pypto.view切出q_v [NPAD, D]k_blk [BY, D]pypto.matmul(q_v, k_blk, pypto.DT_FP32, a_trans=False, b_trans=True)得到[NPAD, BY]分数,经 vector 阶段后pypto.assemblescores_out[0, blk*BY]偏移;
  • host 端 score buffer 也按nb缓存(_scores_cache[NIDX, nb*BY]BF16);
  • 入口带@allow_in_graph@torch.no_grad()

Block 选择逻辑

1. score matmul: idx_q_pad [NPAD, D] @ idx_k [nb*BY, D]^T → scores [NPAD, nb*BY] FP32 2. block max-pool: scores[:NIDX].view(NIDX, nb, BY).amax(-1) → [NIDX, nb] 3. head max: .amax(0) → [nb] block scores 4. top-k: blk[:nb-LOCAL].topk(TOPK-LOCAL) → top-(TOPK-LOCAL) non-local blocks 5. local 强制: arange(nb-LOCAL, nb) → LOCAL 个最近 block 6. concat → [1, TOPK] 短上下文保护: nb <= TOPK 时返回 arange(nb),避免 topk(k) 越界

host 端对应实现(minimax_m3_msa_indexer_impl.py):torch.mm(idx_q, idx_k.t(), out=scores)scores.view(NIDX, nb, BY).amax(dim=(0, 2))(同时完成 block max-pool 与 head max)→blk[:nb-LOCAL].topk(ksel-LOCAL, sorted=False)torch.cat([ids, loc])。短上下文保护与仓库内 modeling 的_msa_decode_block_table保护逻辑一致(当可用的 block 不足TOPK时全部返回,避免topk(k)越界崩溃),见源码注释 minimax_m3_msa_indexer_impl.py。

算子三:minimax_m3_msa_sparse_decode

算法概述

M3 MSA block-sparse decode attention(Q seq-len = 1)。GQA group(Hq // Hkv = 16个 query head 共享 1 个 KV head)batch 进 cube M 轴,对 indexer 选中的 key block 做 online-softmax flash attention,按NTILE宽的 chunk 迭代(比逐 128-block 的 matmul 更大)。当前(部分)block 通过valid_mask列掩码处理。选中的 KV block 在 host 端按 KV head gather 成紧凑张量,block 选择跨 head 共享(因为 M3 indexer 对 index head 做了 max-pool)。

关键常量

常量说明
HQ64query head 数
HKV4KV head 数
GROUP16HQ // HKV,cube M 轴
D128head_dim
BY128KV block size
NTILE512online-softmax chunk 宽度(env:MSA_NTILE
SCALE1/√128attention scale

LARGE_NEG = -3.0e38用于 mask 列置负。NTILE默认 512 是在 sweep 中测得最快的宽度(源码注释 minimax_m3_msa_sparse_attention_impl.py)。

循环结构

outer_loop (B*HKV, static unroll) nchunk_loop (sel // NTILE, static unroll) S = q_block @ k_ch^T * scale [GROUP, NTILE] BF16 → FP32 mask: valid_mask 列掩码(部分 block 置 LARGE_NEG) online softmax: m_c → exp → l_c → p_bf16 O = p_bf16 @ v_ch [GROUP, D] FP32 累加器更新 (mi, li, oi) out = oi / li → cast BF16 → assemble

实现要点(minimax_m3_msa_sparse_attention_impl.py):

  • kernel 按(bsz, topk)形状缓存编译;total_outer = bsz * HKVsel = topk * BYnchunk = sel // NTILE
  • 外层for outer in range(total_outer)静态展开(B*HKV是静态的,利于调度),由b_idx = outer // HKVhkv = outer % HKV推导 batch 与 KV head,q_ofs = b_idx * HQ + hkv * GROUP
  • milioi三个 FP32 累加器用pypto.tensor在片上分配;
  • 每个 chunk:s = matmul(q_block, k_ch, FP32, b_trans=True)mul(SCALE)→ 用mask_row(行广播)做掩码:s = s*mask_row + (1-mask_row)*LARGE_NEGamaxm_cexp(s - m_c)suml_ccast到 BF16 后与v_ch做第二个 matmul(K 按 128 分块);
  • 首 chunk 直接赋值累加器,后续 chunk 做标准 online-softmax 重组:alpha = exp(mi - mi_new)beta = exp(m_c - mi_new)li = alpha*li + beta*l_coi = alpha*oi + beta*o_c
  • 最后out = oi / li,cast 回 BF16 后assemble[q_ofs, 0]

Kernel 签名

minimax_m3_msa_sparse_decode( q, # [B, HQ, D] BF16 — decode query (post norm + RoPE) k_blocks, # [B, HKV, nb, BY, D] BF16 — paged KV cache keys v_blocks, # [B, HKV, nb, BY, D] BF16 — paged KV cache values block_ids, # [B, topk] INT — indexer-selected block ids seq_len, # int — total KV length ) # Returns: [B, HQ, D] BF16 attention output

host 侧(minimax_m3_msa_sparse_attention_impl.py):由seq_len计算cur_block = (seq_len-1)//BYcur_valid,把选中的 block gather 为[HKV, topk*BY, D],先torch.bmm算 score 并乘 scale;若当前 block 是部分块,构造col_valid列布尔掩码(当前 block 的越界列与所有更远的 block 置 False),masked_fill_(-inf)后 softmax,再torch.bmm乘 V。kernel 侧同样接收valid_mask [bsz, sel]作为输入。

性能说明

源码 docstring 诚实记录了性能现状(minimax_m3_msa_sparse_attention_impl.py):对于 M3 decode 形状,原生npu_fused_infer_attention_score(paged block-sparse,经block_table)目前比此 PyPTO kernel快约 5x——这与 MoE grouped-GEMM 的发现一致:对这种 memory-bound、M=16 的 GQA-decode 形状,手调原生 decode kernel 优于 tile-DSL。此 kernel 是 PyPTO-native MSA 路径,原生 op 是性能目标(speed oracle)。源码列出的后续优化方向包括:单次 softmax(去掉 online 重组)、融合 indexer 的 block-max + top-k、以及 on-device gather(gather_in_ub+block_table)。

算子四:msa_main_branch

算法概述

M3 MSA Main Branch:GQA-batched flash attention。16 个 Q head 共享 1 个 KV head,batch 进 cube M 维(per query-block tile),K/V/mask 每个 chunk 加载一次并被 16 个 Q head 复用,大幅降低 HBM 带宽。

循环顺序h_kv → n_block → g(inner chunk) → c(kv chunk),使oi保持片上(单 head 一次),K/V/mask 在外层循环加载并跨所有 16 个 Q head 复用(见 msa_main_branch_impl.py 的模块 docstring)。

参考:MiniMax M3 Technical Report(arXiv:2606.13392v2),Equation 8。

关键常量

常量说明
HQ64query head 数
HKV4KV head 数
GROUP16HQ // HKV
D128head_dim
BK128KV block size
TOPK16selected KV block 数
_MAX_N2048最大 query 序列长度
_NTILE512KV chunk 宽度(env:MSA_NTILE
_NQ_TILE128query-block tile(env:MSA_NQ_TILE

其中_MAX_N可用MSA_MAX_N环境变量或msa_main_branch(..., max_n=...)参数覆盖(如 4096),且max_n必须是bk的整数倍(源码注释见 msa_main_branch_impl.py)。

循环结构

LOOP_HG (total_heads = HKV * GROUP) c_loop (nchunks = kv_len // NTILE) raw = q_all @ k_ch^T [MAX_N, NTILE] → FP32 scaled = raw * scale masked = scaled + mask_ch (causal mask 预计算于 host) online softmax: m_c → exp → l_c → p_cast pv = p_cast @ v_ch [MAX_N, dh] FP32 累加器更新 (mi, li, oi) out = oi / li → assemble

实现要点(msa_main_branch_impl.py):

  • kernel 入口是工厂函数msa_main_branch(hq, hkv, dh, bk, topk, max_n=None),返回一个wrapperwrapper负责把block_mask [num_blocks, topk, bk, bk]transpose + reshape 成[bk*topk, kv_len](必要时用-1e9填充到max_n行),并在_DTYPE != "fp32"时把 Q/K/V cast 到目标 dtype;
  • 外层LOOP_HGpypto.loop(total_heads)遍历HKV * GROUP个 head 组合,由h_kv = hg // groupg = hg - h_kv * group推导 KV head 与 group 内序号;
  • 每个 head:k_all/v_allh_kv_col切列,q_allq_col切列(valid_shape=[n, dh]处理动态 batch);mi/li/oipypto.full分配(valid_shape裁剪到[n, 1]/[n, dh]);
  • 内层for c in range(nchunks)静态展开:raw = matmul(q_all, k_ch, FP32, b_trans=True)scaled = raw * scalemasked = scaled + mask_ch(mask 直接相加,因为 host 已将非因果位置置为大负数)→ online softmax 与累加器更新;
  • p_cast在 dtype 非 FP32 时才 cast(pypto.cast(p, pt_dt) if pt_dt != pypto.DT_FP32 else p)。

Causal Mask

所有 causal 逻辑在 host 端预计算进block_mask [num_blocks, topk, bk, bk],kernel 内无动态条件分支(PyPTOAssignMemoryTypepass 的要求)。掩码规则:

  • kv_seq > qbMASK_NEG(block 不可达)
  • kv_seq == qb:lower-triangular causal mask
  • kv_seq < qb0.0(全注意力,无掩码)

kernel 接收的block_mask[bk*topk, kv_len]FP32(transposed+reshaped,必要时 pad 到max_n行)。

Kernel 签名

msa_main_branch(hq, hkv, dh, bk, topk)(query, key_blocks, value_blocks, block_mask, output) # query: [N, HQ, D] — decode query # key_blocks: [topk*bk, HKV, D] — gathered KV blocks # value_blocks:[topk*bk, HKV, D] — gathered KV blocks # block_mask: [bk*topk, kv_len] FP32 — 预计算 causal mask (transposed+reshaped) # output: [N, HQ, D] FP32

注意 output 是 FP32(不同于其他三个算子),query的 N 维为pypto.DYNAMIC(动态)。

Dtype 支持

通过MSA_DTYPE环境变量选择:fp32(默认)/fp16/bf16_DT_MAP将字符串映射到(pypto.DT_*, torch.*)对(msa_main_branch_impl.py)。

环境变量总览

所有 tile 旋钮和 swigluoai alpha/limit 均可通过环境变量覆盖,汇总如下:

环境变量默认值作用对象
PYPTO_VEC_TILE128(silu)/ 256(swigluoai)Grouped GEMM vector tile 宽度
PYPTO_CUBE_NBUFFER2 / 4Grouped GEMM cube 流水深度
PYPTO_VEC_NBUFFER2 / 1Grouped GEMM vector 缓冲数
PYPTO_L1_REUSE2 / 3Grouped GEMM cube L1 reuse
PYPTO_UNROLL1,2,4,8,16,32,64Grouped GEMM token unroll 列表
PYPTO_MM1_K/PYPTO_MM1_N128 / 256Grouped GEMM mm1 tile
PYPTO_MM2_K/PYPTO_MM2_N128 / 256Grouped GEMM mm2 tile
PYPTO_SWIGLU_ALPHA1.702swigluoai alpha
PYPTO_SWIGLU_LIMIT7.0swigluoai clamp limit
MSA_NTILE512MSA sparse decode / main branch 的 chunk 宽度
MSA_NQ_TILE128main branch query-block tile
MSA_MAX_N2048main branch 最大 query 序列长度
MSA_DTYPEfp32main branch 计算/存储 dtype(fp32/fp16/bf16)
USE_PTO_GROUPED_GEMMFalse(代码内开关)是否将 MoE FFN 路由到融合 kernel

此外,测试运行依赖TILE_FWK_DEVICE_ID(默认 0)指定 NPU 设备。

测试用例与精度验证

minimax_moe_grouped_gemm(M2.7 / silu)

用例来自 tests/ops/minimax_m27/test_cases.json:

用例EHIcounts说明
case_001830721536[0,1,2,4,8,16,32,0]混合 token 数(宽度 1/2/4/8/16/32),含零 token expert;seed=321,rtol=atol=0.008
export TILE_FWK_DEVICE_ID=0 python3 tests/ops/minimax_m27/test_minimax_m27_grouped_gemm.py # 或 pytest python3 -m pytest tests/ops/minimax_m27/test_minimax_m27_grouped_gemm.py

测试脚本还支持--list列出用例与按case_id单跑。golden 参考(test_minimax_m27_grouped_gemm.py)逐 expert 用functional.linear+functional.silu独立计算,与 kernel 输出做numpy.testing.assert_allclose

minimax_moe_grouped_gemm(M3 / swigluoai)

用例来自 tests/ops/minimax_m3/test_cases.json:

用例EHIcounts说明
case_smoke_edge4256128[1,0,3,4]不均匀路由(含零 token expert)+ swigluoai clamp 覆盖;seed=99,assert_clamp=true,rtol=0.015 / atol=0.5
python3 tests/ops/minimax_m3/test_minimax_m3_grouped_gemm.py python3 -m pytest tests/ops/minimax_m3/test_minimax_m3_grouped_gemm.py

该用例特意用input_scale=0.5/weight_scale=0.3放大数据,并断言 gate/up 确实越过 ±limit,确保 clamp 分支被执行到(见 test_minimax_m3_grouped_gemm.py)。M3 golden 独立实现了 swigluoai 公式并在 activation 输出处模拟 kernel 的 BF16 舍入。

minimax_m3_msa_indexer + sparse_decode

用例来自 tests/ops/minimax_m3/test_minimax_m3_msa_pypto.py:

用例nb说明
test_indexer_selection_identical_to_torch[8]8nb≤TOPK 短上下文保护
test_indexer_selection_identical_to_torch[17]17最短 sparse 路径(TOPK+1)
test_e2e_indexer_plus_attention17indexer→attention 端到端(skip,需空闲 die)
test_attention_forward_pypto_msa_matches_native_paged_decodeHF attention 集成(skip,需完整 HF modeling 环境)
python3 -m pytest tests/ops/minimax_m3/test_minimax_m3_msa_pypto.py -v

精度口径:indexer 选中的 block id集合与 torch reference 完全一致(pypto_ids == torch_ids);e2e 用例的 max_diff < 5e-2。这些用例在无 Ascend 环境时会因缺torch_npu直接报错(NPU-only)。

msa_main_branch

配置
HQ64
HKV4
TOPK16
N2048
HEAD_DIM128
BK128
export TILE_FWK_DEVICE_ID=0 python3 tests/ops/minimax_m3/test_msa.py # 带泳道图采集 COLLECT_SWIMLANE=1 python3 tests/ops/minimax_m3/test_msa.py # msprof 性能采集 python3 tests/ops/minimax_m3/profile_msa.py python3 tests/ops/minimax_m3/profile_golden_msa.py

精度校验口径汇总

  • Grouped GEMM:numpy.testing.assert_allclose,rtol=0.008~0.015,atol=0.008~0.5(视变体和 clamp 覆盖而定;M3 因 clamp 处的 BF16 舍入放宽到 0.5)
  • MSA indexer:选中的 block id 集合与 torch reference完全一致(identical)
  • MSA sparse decode / main branch:max_diff < 5e-2 / atol_abs=1e-3, atol_rel=1e-3

运行方式

# 设置设备 ID export TILE_FWK_DEVICE_ID=0 # M2.7 Grouped GEMM (silu) python3 -m pytest tests/ops/minimax_m27/ # M3 Grouped GEMM (swigluoai) + MSA indexer/decode python3 -m pytest tests/ops/minimax_m3/ # MSA Main Branch python3 tests/ops/minimax_m3/test_msa.py

与模型迁移的衔接

算子层之上,modeling/transformers/minimax/README.md 记录了完整的 NPU 迁移链路:

  • M2.7(hidden_size=3072I=1536E=256、top-8、SiLU-SwiGLU)与 M3(hidden_size=6144I=3072E=128、top-4、clamped GLU)共用一套 runner,由--variant {m27,m3}选择;
  • MoE expert FFN 路由到共享的 PyPTO 融合 grouped-GEMM 算子(routed experts 在 host 保持 FP8,逐层反量化后流式送入 BF16 kernel);M3 的 MSA decode 路径在--use_pypto下启用;
  • 实测/精度验证注意:M3 的 layer 0-2 为 dense、layer 3+ 为 MoE,因此 benchmark 需--max-layers >= 4才会真正触发 grouped-GEMM 路径。

源码归档映射(README 的 Archive Mapping 一节)也印证了本目录的角色:minimax_m27/m3的 modeling 文件落入src/pypto_gym/transformers/,而 grouped-GEMM / MSA 算子即本目录src/pypto_gym/ops/pypto_tensor/minimax/,算子测试对应tests/ops/minimax_m27/tests/ops/minimax_m3/

小结

本文围绕 MiniMax 算子 README 系统梳理了四个 PyPTO 融合算子的算法、tile 设计与验证方法:minimax_moe_grouped_gemm以单一 kernel 覆盖 M2.7/M3 两个 MoE 变体(差异仅在 activation 与 tile 默认值),minimax_m3_msa_indexerminimax_m3_msa_sparse_decode组成 M3 sparse decode 的 PyPTO-native 路径,msa_main_branch通过 GQA batching 将 16 个 Q head 复用同一份 K/V 以削减 HBM 带宽。所有旋钮均可通过PYPTO_*/MSA_*环境变量在不改源码的前提下重调,测试与 golden 参考可直接在 tests/ops/minimax_m27/ 与 tests/ops/minimax_m3/ 中复现。对于追求极致 decode 性能的场景,源码也如实指出了原生npu_fused_infer_attention_score目前更快的现状,PyPTO-native 路径的定位是"结构等价 + 可调优的 DSL 实现",后续优化方向(单次 softmax、融合 indexer 尾巴、on-device gather)已记录在实现源码的 docstring 中。

【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询