ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

LLaDA2.0-mini MoE 融合算子实战:基于 PyPTO 的 Grouped GEMM 单内核实现与 NPU 集成指南

2026/9/18 23:22:48 拓冰建站 浏览量
LLaDA2.0-mini MoE 融合算子实战:基于 PyPTO 的 Grouped GEMM 单内核实现与 NPU 集成指南 LLaDA2.0-mini MoE 融合算子实战基于 PyPTO 的 Grouped GEMM 单内核实现与 NPU 集成指南【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读本文以 CANN pypto-gym 仓库中的 LLaDA2.0 MoE fused operators 文档 为核心系统讲解针对 LLaDA2.0-miniinclusionAI 的扩散式语言模型 MoE 架构的 PyPTO 融合算子实现——llada2_moe_grouped_gemm。该算子把LLaDA2MoeSparseMoeBlock中原本逐专家per-expert的 Python 循环推理合并为一次覆盖全部专家的 Grouped GEMM 内核调用。读完本文你将掌握该内核的输入布局、JIT 配置、SwiGLU 融合计算细节、USE_PTO_EXPERT_FFN开关的接入方式以及如何在昇腾 NPU 上运行算子级测试与端到端推理/基准验证。一、背景为什么 MoE 专家 FFN 需要融合算子LLaDA2.0-mini 是一个 Diffusion-LM 混合专家Mixture-of-Experts模型其 MoE 块包含大量专家LLaDA2.0-mini 的典型规模为H2048, I512, E256即隐藏维度 2048、中间维度 512、专家数 256。每个专家本质是一个 SwiGLU 前馈网络gate/up/down 三个线性层结构对应modeling_llada2_moe.py中的LLaDA2MoeMLPclass LLaDA2MoeMLP(nn.Module): def forward(self, x): return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))原始的moe_infer推理路径见moe_infer会先按路由结果对 token 排序然后对每个专家逐一执行expert(sorted_tokens[start:end])最后拼接并逆排序。当专家数达到 256 时这种 per-expert 的 Python 调度循环本身的开销会显著放大kernel launch 延迟、host 端循环、碎片化的小矩阵乘。这正是 Grouped GEMM 融合的动机把多个专家、多次小 GEMM压缩为一次内核调用、内部按专家切分。二、算子总览与产品支持2.1 算子清单Operator实现文件说明llada2_moe_grouped_gemmllada2_moe_grouped_gemm_impl.py全部专家共享一次 Grouped GEMM 调用替代 per-expert 调度循环算子输入输出统一为BF16FP32 累加与模型推理精度要求一致。2.2 产品与硬件支持产品/系列支持状态Ascend 950PR不支持Atlas A3 训练/推理系列支持Atlas A2 训练/推理系列支持模型侧迁移文档modeling/transformers/llada2_moe/README.md给出了仓库中验证过的软件环境torch_npu 2.10、transformers 4.57.1、pypto 0.2.1、pto-isa v9.1.0、CANN 9.1.0、NPU 为 Ascend 910B3。算子级测试还要求环境已安装torch_npu与pypto见下文测试章节。三、内核设计一次调用完成所有专家的 SwiGLU FFN3.1 输入布局由 host 侧准备Grouped GEMM 内核接收已经按专家排好序的 token 与按行拼接的权重全部为 ND 格式pypto.TileOpFormat.TILEOP_ND参数形状类型含义sorted_tokens[N_total, H]BF16按专家分配预排序后的 tokenw13_flat[E*H, 2*I]BF16所有专家的 gate‖up 权重按行拼接w2_flat[E*I, H]BF16所有专家的 down 权重按行拼接expert_cumsum[E1]INT32专家累计 token 数专家 e 拥有[cumsum[e], cumsum[e1])result[N_total, H]BF16输出缓冲与sorted_tokens同形状非张量参数num_experts专家数 E、hidden_sizeH、intermediate_sizeI。3.2 双层循环pypto.loop×pypto.loop_unroll内核核心结构对应llada2_moe_grouped_gemm_kernel外层用pypto.loop(0, num_experts, 1, nameEXPERT_LOOP)遍历专家维度内层用pypto.loop_unroll(0, n_e, 1, nameLOOP_TOKEN, unroll_list[1, 2, 4, 8, 16, 32, 64])按 token 维度展开。展开列表支持 164 的 2 的幂次从而以动态 tile 尺寸覆盖每个专家分到的 token 数变化在 unroll 循环内部通过pypto.set_cube_tile_shapes设置 GEMM 的 tile 形状再以pypto.view从扁平权重中切出当前专家的w13_e w13_flat[e_idx*H : (e_idx1)*H, :]与w2_e w2_flat[e_idx*I : (e_idx1)*I, :]并从sorted_tokens切出当前专家的 token 块tile_x。每个专家内部的计算流是标准的两阶段 SwiGLU FFNgate_up tile_x w13_e (FP32 累加输出 2*I 宽) sw SiLU(gate) * up (FP32 计算转 BF16) down sw w2_e (FP32 累加) out cast(down, BF16)随后pypto.assemble(out, [e_start tok_idx, 0], result)将结果写回result的对应行区间。3.3 SwiGLU 的 PyPTO 实现_swiglu_silu实现代码把gate_up按中间维度一分为二前半为 gate、后半为 up然后依次用pypto.mul/pypto.exp/pypto.add/pypto.div/pypto.mul组合出 SiLU 激活def _swiglu_silu(gate_up): half gate_up.shape[1] // 2 gate pypto.view(gate_up, [gate_up.shape[0], half], [0, 0]) up pypto.view(gate_up, [gate_up.shape[0], half], [0, half]) neg pypto.mul(gate, -1.0) e pypto.exp(neg) denom pypto.add(e, 1.0) silu pypto.div(gate, denom) return pypto.mul(silu, up)注意gate_up的 matmul 输出类型是pypto.DT_FP32即GEMM 累加在 FP32 中进行SwiGLU 激活也在 FP32 完成之后才 cast 回 BF16 参与第二个 GEMM最后再 cast 回 BF16 写出。这与模型推理中BF16 输入输出、FP32 累加的精度策略完全一致。四、JIT 运行时配置与可调环境变量4.1pypto.frontend.jit选项内核通过pypto.frontend.jit装饰器声明编译选项源码位置runtime_optionsdevice_sched_mode: 1设备调度模式stitch_function_max_num: 64最大算子缝合stitch函数数配合多次 matmul vector 指令融合ready_on_host_tensors: [expert_cumsum]声明expert_cumsum需要在 host 端就绪用于按专家切片源码注释说明这是 CANN 9.1.0 下的必要配置。pass_optionscube_l1_reuse_setting: {-1: 2}Cube 单元 L1 复用设置cube_nbuffer_setting: {-1: _CUBE_NBUF}与vec_nbuffer_setting: {-1: _VEC_NBUF}Cube/Vector 多缓冲双缓冲设置。此外内核还调用pypto.experimental.set_operation_options(combine_axisTrue)开启轴合并优化。4.2 可调环境变量实现文件顶部从环境变量读取默认值用于覆盖 tile 尺寸与缓冲配置环境变量默认值作用PYPTO_CUBE_NBUFFER2Cube 侧多缓冲数量PYPTO_VEC_NBUFFER2Vector 侧多缓冲数量PYPTO_MM1_M8第一个 GEMMtile_x w13的 M 维度 tilePYPTO_MM1_K128第一个 GEMM 的 K 维度 tilePYPTO_MM1_N256第一个 GEMM 的 N 维度 tilePYPTO_MM2_M8第二个 GEMMsw w2的 M 维度 tilePYPTO_MM2_K128第二个 GEMM 的 K 维度 tilePYPTO_MM2_N256第二个 GEMM 的 N 维度 tilePYPTO_VEC_FIRST13Vector 阶段首次 tile 的 token 数上限min(tile_batch, vec_first)这些参数在运行内核前以export PYPTO_XXX...方式注入便于在不改动代码的情况下针对具体形状做调优。内核中pypto.set_cube_tile_shapes的第一个参数为[tile_batch, tile_batch]M 维 N 维 token tile第二、三个参数分别为 K/N 维 tile 组合如[mm1_cube[1], mm1_cube[1]*2]对应 K128、输出 2*I 宽度。4.3 host 封装与输入校验对外暴露的llada2_moe_grouped_gemm源码具备两个关键特性用torch._dynamo.allow_in_graph装饰可被torch.compile/ Dynamo 图捕获若sorted_tokens是FakeTensor编译期 shape 推导阶段直接返回而不实际执行通过_check做输入形状/类型校验sorted_tokens、w13_flat、w2_flat必须是 BF16 二维张量expert_cumsum必须是[E1]的 INT32 一维张量w13_flat形状必须等于(E*H, 2*I)、w2_flat等于(E*I, H)result与sorted_tokens同形状。五、模型侧集成_moe_infer_pypto与USE_PTO_EXPERT_FFN开关5.1 集成位置内核通过LLaDA2MoeSparseMoeBlock._moe_infer_pypto()接入模型源码由USE_PTO_EXPERT_FFN开关门控opt-in默认关闭。开关定义在算子适配模块llada2_moe/__init__.py中from .llada2_moe_grouped_gemm_impl import llada2_moe_grouped_gemm as grouped_gemm USE_PTO_EXPERT_FFN False运行时会话通过sys.modules[llada2_pto_kernels]暴露该模块moe_infer检测到USE_PTO_EXPERT_FFN为 True 时改走_moe_infer_pypto见moe_infer中的PYPTO_PATCH分支。5.2 PyPTO 推理路径的处理流程_moe_infer_pypto在 host 侧完成以下编排权重合并CPU 中转控峰值内存_ensure_pypto_weights把每个专家的gate_proj/up_proj/down_proj权重转置后按行堆叠成_pypto_w13_stack [E, H, 2I]与_pypto_w2_stack [E, I, H]。为了控制显存峰值大 E 下原件 堆叠件同时驻留 NPU 会在 64GB die 上 OOM实现先在 CPU 上暂存堆叠结果拷贝完每个专家后立即释放其设备端权重最后再把完整堆叠搬运回设备并注册为非持久 bufferpersistentFalse。token 排序与 cumsumflat_ids.argsort()得到按专家排序的置换torch.bincount统计各专家 token 数torch.cumsum生成expert_cumsum [E1]INT32。权重展平reshape(E*H, 2I)与reshape(E*I, H)得到内核所需的扁平布局。越界保护pad由于动态偏移下 MTE 可能对排序后的输入/输出缓冲多读写至多一个最大 tilecodegen 预取/对齐产物实现按最大 tilepad_rows 64对输入输出做F.pad避免触发 aicore 错误 507015The DDR address of the MTE instruction is out of range内核执行后截断回有效行数。逆排序与加权求和inv[sort_perm] arange建立逆置换把 Grouped GEMM 的输出还原到原始 token 顺序按topk_weight加权求和得到最终 MoE 输出。集成范围来自 src/pypto_gym/transformers/llada2_moe/README.md中目前只有Expert FFNGrouped GEMM走 PyPTORouter gating、RMSNorm、AttentionGQA、Dense MLP、RoPE 均保留 PyTorch 原实现作为回退。5.3 运行端到端推理与基准模型仓库侧的 READMEmodeling/transformers/llada2_moe/README.md给出了完整流程export MODEL_PATH/path/to/LLaDA2.0-mini # 下载权重 python3 modeling/transformers/download_hf_model.py \ --model-id inclusionAI/LLaDA2.0-mini \ --output-dir $MODEL_PATH # 将 checkpoint patch 到仓库内模型定义 python3 modeling/transformers/runtime_patch.py \ --model-family llada2_moe \ --model-path $MODEL_PATH # 基线推理 python3 modeling/transformers/llada2_moe/ask_LLaDA2-mini.py --model-path $MODEL_PATH --device NPU # 开启 PyPTO 融合算子推理 python3 modeling/transformers/llada2_moe/ask_LLaDA2-mini.py --model-path $MODEL_PATH --device NPU --use_pypto # 端到端基准基线 vs PyPTO MODEL_PATH$MODEL_PATH DEVICE0 bash modeling/transformers/llada2_moe/bench_LLaDA2-mini.shask_LLaDA2-mini.py在--use_pypto时把仓库内的modeling_llada2_moe.py安装到模型目录备份原文件为.orig导入pypto_gym.ops.pypto_tensor.llada2_moe模块并把USE_PTO_EXPERT_FFN置 True、注册到sys.modules[llada2_pto_kernels]脚本逻辑。注意 LLaDA2 采用 block-wise masked diffusion 生成generate()传的是gen_length/steps/block_length而非max_new_tokens。基准脚本bench_LLaDA2-mini.sh先跑 warmup 吸收 JIT 编译开销再分别测量基线与 PyPTO 模式输出time_mean_s、tps_mean、generate_peak_mem_mb三项对比。六、算子级测试与正确性验证README 给出的测试命令export TILE_FWK_DEVICE_ID0 python3 -m pytest tests/ops/llada2_moe/测试用例来自模型真实形状/精度balanced / uneven / zero-token / mixed-width测试入口为 tests/ops/llada2_moe/test_llada2_moe_grouped_gemm.py。核心验证逻辑测试要求 Ascend NPU 环境import torch_npu失败即报RuntimeError测试前开启torch_npu.npu.config.allow_internal_format True设备 ID 取自TILE_FWK_DEVICE_ID默认 0以H2048, I512为基准构造 token 计数[0, 1, 2, 4, 8, 16, 32, 0]覆盖零 token 专家、非 2 的幂 token 数、奇数/偶数专家分布等边界场景参照实现reference_per_expert用纯 PyTorch 逐专家x w13→F.silu(gate) * up→sw w2计算FP32 累加与 Grouped GEMM 内核输出做np.testing.assert_allclose(rtol8e-3, atol8e-3)比对测试中随机数种子固定为 42权重与 token 按* 0.02缩放以贴近真实模型的小数值分布。测试同时验证了 host 侧参数校验路径_check与FakeTensor短路行为可作为理解内核调用契约的补充参考。七、从源码结构看可复用模式扁平权重 cumsum 索引的 Grouped GEMM 布局是本算子的核心抽象w13_flat/w2_flat按专家行拼接、expert_cumsum界定每个专家 token 区间这种布局在 pypto-gym 的同类算子如 minimax_grouped_gemm_impl.py中也有体现可作为多专家模型融合的通用范式。allow_in_graphFakeTensor短路的封装方式使算子能够平滑嵌入torch.compile图与 shape 推导流程而不会在编译期真正下发 NPU 任务。运行时开关 sys.modules注册的 opt-in 集成策略USE_PTO_EXPERT_FFN默认 False让融合路径与基线路径并存便于逐算子灰度验证。八、使用前提与限制硬件上要求 Atlas A2 / A3 训练或推理系列 NPUAscend 950PR 不支持并安装与 CANN 9.1.0 配套的torch_npu、pypto、pto-isa环境算子当前聚焦推理路径moe_infer/_moe_infer_pypto训练侧forward的 per-expert 循环self.training分支仍保留 PyTorch 实现PyPTO 路径仅覆盖 Expert FFN路由器、归一化、注意力等仍走 host 侧 PyTorch因此端到端加速收益主要来自专家 GEMM 的调度与融合开销削减在真实稀疏路由下内核因动态偏移存在至多一个最大 tile 的越界读写风险host 侧必须按pad_rows64对输入输出缓冲做 padding该约束已内建于_moe_infer_pypto直接调用底层llada2_moe_grouped_gemm时需自行处理。九、关键文件索引文件作用src/pypto_gym/ops/pypto_tensor/llada2_moe/README.md本文主题文档算子概述、产品支持、测试与集成说明src/pypto_gym/ops/pypto_tensor/llada2_moe/llada2_moe_grouped_gemm_impl.pyGrouped GEMM 内核实现JIT 配置、双层循环、SwiGLU、host 封装src/pypto_gym/ops/pypto_tensor/llada2_moe/init.py算子适配模块与USE_PTO_EXPERT_FFN开关定义src/pypto_gym/transformers/llada2_moe/modeling_llada2_moe.py模型图_moe_infer_pypto()与_ensure_pypto_weights()集成逻辑src/pypto_gym/transformers/llada2_moe/configuration_llada2_moe.pyLLaDA2MoeConfig专家/分组拓扑与路由参数tests/ops/llada2_moe/test_llada2_moe_grouped_gemm.py数值正确性测试混合 token 计数、BF16 容差比对modeling/transformers/llada2_moe/README.md端到端 NPU 迁移环境与使用步骤modeling/transformers/llada2_moe/ask_LLaDA2-mini.py基线/PyPTO 双模式推理脚本modeling/transformers/llada2_moe/bench_LLaDA2-mini.sh基线 vs PyPTO 端到端基准脚本【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考