ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyPTO Gym 中的 grouped_matmul_swiglu_quant:MXFP8 分组矩阵乘法 + SwiGLU 激活 + Per-Token 量化的端到端融合算子

PyPTO Gym 中的 grouped_matmul_swiglu_quant:MXFP8 分组矩阵乘法 + SwiGLU 激活 + Per-Token 量化的端到端融合算子 PyPTO Gym 中的 grouped_matmul_swiglu_quantMXFP8 分组矩阵乘法 SwiGLU 激活 Per-Token 量化的端到端融合算子【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读grouped_matmul_swiglu_quant是 PyPTO-Gym 仓库中面向 MoEMixture of Experts推理场景的高密度融合算子将「MXFP8 分组矩阵乘法Grouped MatMul→ SwiGLU 激活 → Per-Token INT8 动态量化」三步计算合并到单个 PyPTO kernel 中执行。本文以算子 README 为骨架结合仓库内的 PyPTO 实现、纯 PyTorch Golden 参考实现与端到端单测系统讲解该算子的数学语义、输入输出规格、Shape 约束、显式 tile 配置、量化流水线及精度验证方法。读完本文你将掌握如何在 PyPTO 新前端 JIT 下用 Cube Vector 融合方式落地一个 MoE 后处理量化算子并了解其与gmm_mxfp8等相邻算子的演进关系。算子定位MoE 场景中的后处理融合grouped_matmul_swiglu_quant位于仓库的 src/pypto_gym/ops/pypto_tensor/experimental/matmul/grouped_matmul_swiglu_quant/ 目录是 MoE 场景中 grouped matmul 的后处理融合算子。它完成三件事MXFP8 grouped matmul按group_list将路由后的 token 切分给不同 expert每个 expert 使用独立权重执行带缩放因子的矩阵乘法SwiGLU 激活将 matmul 输出沿最后一维二等分为 value/gate 两半计算value * sigmoid(value) * gatePer-token INT8 动态量化按每个 token 独立求最大绝对值生成 scale将 SwiGLU 结果量化到 INT8同时输出每个 token 对应的 FP32 scale。产品支持情况算子 README 明确声明了当前支持矩阵三列均为不支持说明该算子处于实验验证阶段Ascend 950PR不支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持需要说明的是仓库单测 test_gmm_swiglu_quant.py 中带有pytest.mark.soc(950)标记表明测试用例是为 950 架构准备的在未开放支持的硬件上运行前应以 README 的支持矩阵为准确认可用性。算子语义与数学公式数学公式对第 i 个 expertREADME 给出了完整的量化公式链gmm_out_i ScaledMatmul(a_i, b_i, scaled_a_i, scaled_b_i) value_i, gate_i chunk(gmm_out_i, 2, dim-1) swiglu_i value_i * sigmoid(value_i) * gate_i scale_i max(abs(swiglu_i), dim-1, keepdimTrue) / 127 out_i clamp(round(swiglu_i / scale_i), -127, 127).to(int8) out_quant_i scale_i.squeeze(-1)展开形式按元素表达gmm_out[m, n] Σ(k0..K-1) dequant(a[m, k]) * dequant(b[expert(m), k, n]) output[m, d] Quant(SiLU(gmm_out[m, d]) * gmm_out[m, d N/2])其中dequant由 MXFP8 输入值和 E8M0FNU scale 共同决定。这里体现了 MXFP8Microscaling FP8量化格式的核心思想每 64 个元素共享一个仅含指数部分的 E8M0FNU 缩放因子数据部分为 FP8 E4M3 格式dequant(x) x * scale从而以极低开销完成高精度浮点的近似表达。四步计算流程README 将整个算子拆解为四个阶段Grouped Matmul 计算按group_list切分 token逐 expert 调用pypto.scaled_mm。a_i: [M_i, K]b_i: [K, N]或[N, K]输出gmm_out_i: [M_i, N]SwiGLU 激活将gmm_out_i沿最后一维二等分。value: [M_i, N/2]gate: [M_i, N/2]swiglu value * sigmoid(value) * gatePer-token 量化按 token 求最大绝对值并生成 scale。scale max(abs(swiglu), dim-1) / 127output round(swiglu / scale).to(int8)输出组装将每个 expert 的 INT8 输出和 FP32 scale 组装回全局输出。输入输出规格输入张量名称ShapeDType说明a[M, K]FP8 E4M3路由后 token 输入b[E, K, N]或[E, N, K]FP8 E4M3expert 权重scaled_a[M, K/64, 2]E8M0FNUtoken scale每 64 元素共享一个scaled_b[E, K/64, N, 2]或[E, N, K/64, 2]E8M0FNUexpert 权重 scalegroup_list[E]int/list每个 expert 的 token 数注意scaled_a/scaled_b的最后一维为 2MXFP8 每 64 个数据元素对应一组 scale 分量README 以K/64表达 scale 的块粒度实际 shape 的最后一个维度2承载 scale 的两部分对应 MXFP8 每 64 元素 scale 组的结构在测试构造时以torch.float8_e8m0fnu生成见 test_gmm_swiglu_quant.py。输出张量名称ShapeDType说明out[M, N/2]int8SwiGLU 后 per-token 量化输出out_quant[M]float32每个 token 对应的量化 scale由于 SwiGLU 把 N 维折叠为 N/2out的列数只有 matmul 输出列数的一半out_quant则是完全 per-token 粒度的 scale 向量供下游反量化使用。Shape 范围与约束动态轴轴当前覆盖范围说明M16路由后 token 总数K512matmul K 维N7168matmul 输出列数必须为偶数E2expert 数量约束条件N 必须为偶数SwiGLU 需要将最后一维切分为 value/gate 两半group_list 和 M 一致sum(group_list) M当前测试覆盖 b_transFalse代码保留 transpose 配置内置 case 使用非转置权重布局K 与 scale 布局匹配当前 scale 按K / 64构造输出 scale 为 per-token 粒度out_quantshape 为[M]。实现特点与源码级剖析算子的 PyPTO 实现位于 gmm_swiglu_quant_impl.py以下逐条对应 README 的实现特点展开。1. 新前端 JITkernel 使用pypto.frontend.jit装饰并显式传入运行期选项pypto.frontend.jit( runtime_options{device_sched_mode: 3}, ) def scaled_matmul_kernel(a, b, scaled_a, scaled_b, out, out_quant, group_list, tile_config): ...device_sched_mode: 3用于指定设备调度模式属于 kernel 级运行期配置可在不改动计算逻辑的前提下影响任务在 NPU 上的调度方式。2. Cube Vector 融合pypto.scaled_mm在 Cube 单元上完成 MXFP8 缩放矩阵乘并输出 FP32SwiGLU 与量化全程使用 Vector 单元指令sigmoid、mul、cast、abs、amax、div、round实现了 Cube 与 Vector 两类计算单元的流水融合current_mm_out pypto.scaled_mm(x, weight, pypto.DT_FP32, scaled_x, scaled_weight) value current_mm_out[:, : n_size // 2] gate current_mm_out[:, n_size // 2:] silu_value value * pypto.sigmoid(value) swiglu_out pypto.mul(silu_value, gate)3. 按 expert 分段处理kernel 通过group_list维护滚动窗口begin/end逐 expert 切片输入并调用scaled_mm保持 grouped matmul 语义for i in range(num_groups): begin end end end group_list[i] x a[begin:end, :] weight b[i] scaled_x scaled_a[begin:end, :, :] scaled_weight scaled_b[i] ...4. 显式 tile 配置每个 expert 的 matmul 之前通过pypto.set_cube_tile_shapes显式配置 Cube 分块并开启 SplitKSwiGLU/量化阶段通过pypto.set_vec_tile_shapes配置 Vector tilepypto.set_cube_tile_shapes( tile_config.m_tile_shape, tile_config.k_tile_shape, tile_config.n_tile_shape, enable_split_kTrue, ) ... pypto.set_vec_tile_shapes(64, 256)从单测 test_gmm_swiglu_quant.py 可以看到内置 case 的实际 tile 参数参数值tile_size256m_tile_shape[9, 9]k_tile_shape[256, 256]n_tile_shape[256, 256]vector_tile_shape[1, 8, 256, 32]这些参数封装在ShapeConfigdataclass 中同时携带a_trans/b_trans/*_format_nz等布局标志通过tile_config传入 kernel便于针对不同 Shape 快速切换分块策略。5. 量化流水线的数值细节实现中的 per-token 量化并非一步到位而是采用「BF16 圆整 → FP32 放大 → 取 abs/max → 求 scale → 乘 scale → round → FP16 → INT8」的多级数值链x_bf16 pypto.cast(swiglu_out, pypto.DT_BF16, pypto.CastMode.CAST_RINT) x_fp32 pypto.cast(x_bf16, pypto.DT_FP32) x_abs pypto.abs(x_fp32) x_max pypto.amax(x_abs, -1, True) x_scale pypto.div(pypto.full([shape_0, shape_1], 127.0, pypto.DT_FP32), x_max) x_mul pypto.mul(x_fp32, x_scale) x_mul_round pypto.round(x_mul) x_fp16 pypto.cast(x_mul_round, pypto.DT_FP16, pypto.CastMode.CAST_RINT) x_int8 pypto.cast(x_fp16, pypto.DT_INT8) x_scale_quant pypto.div(pypto.full([shape_0, shape_1], 1.0, pypto.DT_FP32), x_scale)其中CAST_RINT指定了向最近整数舍入的转换模式x_scale 127 / max与 README 中scale max/127互为倒数关系最后通过x_scale_quant 1 / x_scale还原回 README 约定的max/127语义再写入out_quant。这种「除 127 → 乘 127」的写法是出于数值稳定与指令开销的权衡。6. 内存访问模式与输出组装README 总结的内存访问模式在源码中均有对应a和scaled_a按 expert token 范围连续切片a[begin:end, :]b和scaled_b按 expert 维度读取b[i]、scaled_b[i]out和out_quant使用pypto.assemble按 token 起始偏移写回pypto.assemble(x_int8, [begin, 0], out) pypto.assemble(x_scale_quant, [begin, 0], out_quant)7. Host 侧封装gen_mxfp8是 host 侧入口将输入张量搬运到 NPU.npu()、预分配outINT8与out_quantFP32输出 buffer、调用 kernel最后把输出还原为约定的 shapeout torch.zeros((a.shape[0], b.shape[-1] // 2), dtypetorch.int8).npu() out_quant torch.zeros((a.shape[0], 1), dtypetorch.float32).npu() scaled_matmul_kernel(a, b, scaled_a, scaled_b, out, out_quant, inputs.group_list, tile_config) out out.to(torch.float32) out_quant out_quant.squeeze(dim1)注意 kernel 内部实际写入的out_quantshape 为[M, 1]host 侧在 kernel 返回后squeeze(dim1)得到 README 约定的[M]。精度验证容差设置对比对象相对容差 (RTOL)绝对容差 (ATOL)INT8 输出0.0011Scaleout_quant0.00010.0001INT8 输出的 ATOL 放宽到 1是因为量化取整本身存在半格误差scale 属于连续量容差收紧到 1e-4。测试用例测试名称MKNgroup_list说明testcase6165127168[7, 9]非均匀 2 expert grouped matmul SwiGLU quant测试代码中该 case 的注册名同时保留为testcase参数与testcase6完全一致见 get_params。[7, 9]是一个非均匀切分用于覆盖「两个 expert 分到的 token 数不同」这一 MoE 真实场景。验证方法Golden 实现gmm_swiglu_quant_golden.py 中gen_golden为纯 PyTorch 参考实现PyPTO 实现gmm_swiglu_quant_impl.py 中gen_mxfp8调用scaled_matmul_kernel对比工具numpy.testing.assert_allclose。Golden 参考实现要点Golden 侧需要模拟 MXFP8 反量化语义其核心逻辑compute_golden_result包括按transpose.a_trans / b_trans对输入/权重及 scale 做转置与 reshape 对齐处理ceil(K/32)为奇数时的 scale 裁剪scaled_x_golden[:, :-1]、scaled_weight_golden[:-1, :]将 per-64 的 scale 通过repeat_interleave(..., repeats32, ...)广播到每个元素并对 K 维不足部分做 zero-padding先反量化FP8 值 × scale转 FP32再执行torch.matmulswiglu()与quant_pertoken()分别复现激活与量化torch.chunk(2, dim-1)、torch.round、clamp(-127, 127)、to(torch.int8)gen_golden按 group_list 逐 expert 计算后用torch.cat拼接全局输出。值得一提的细节Golden 中 SwiGLU 输出先转bfloat16再转回float32用于模拟 PyPTO 实现中「FP32 → BF16 → FP32」的中间精度转换从而让两侧数值路径保持一致。端到端测试流程测试 test_gmm_mxfp8 依次完成构造 tile 配置 → 构造随机测试张量FP8 E4M3 输入、E8M0FNU scale→ 运行 Golden → 运行 PyPTO kernel → 对out与out_quant分别执行assert_allclose。运行单测README 给出两种运行方式需在单测目录或配置好PYTHONPATHcd tests/ops/experimental/matmul/grouped_matmul_swiglu_quant PYTHONPATH/path/to/pypto-gym/src python test_gmm_swiglu_quant.py python test_gmm_swiglu_quant.py testcase6从 test_gmm_swiglu_quant.py 的头部逻辑看脚本自身会向上查找包含src的目录并注入PYTHONPATHsrc与src/pypto_gym/ops/pypto_tensor都会被加入sys.path因此在实际仓库中直接python test_gmm_swiglu_quant.py即可运行无需手动设置环境变量。与相邻算子的关系grouped_matmul_swiglu_quant可以视为 gmm_mxfp8 的功能超集后者仅实现 MXFP8 分组矩阵乘法并输出 FP32 结果前者在相同scaled_mmgroup_list分组框架之上进一步融合了 SwiGLU 激活与 per-token INT8 动态量化将「反量化 → 激活 → 再量化」链路完整收入 kernel 内部避免中间结果落回 host。两者共享TransposeConfig、ShapeConfig等 dataclass 设计与b_trans布局约定可以从 matmul 目录总览 对比阅读理解 PyPTO 实验性 matmul 算子族「从纯 GMM 到融合后处理量化」的演进脉络。小结grouped_matmul_swiglu_quant是一个完整的 PyPTO 融合算子教学案例它以 MoE 推理中最常见的「expert 分组 matmul SwiGLU 量化」三段式计算为场景展示了新前端 JIT、Cube/Vector 显式分块、SplitK 开启、scaled_mm复用、assemble按偏移写回等 PyPTO 编程要点并配套了数值语义对齐的纯 PyTorch Golden 与端到端assert_allclose校验。读者在复用或扩展该算子时应重点核对 README 中的支持矩阵、N 为偶数约束、sum(group_list) M一致性以及 K 与 scale 布局的匹配关系这些约束共同决定了算子在新的 Shape 组合下能否正确运行。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表