CANN pypto-gym 实战:transpose_quant_batch_matmul 算子的 MXFP8 量化批量矩阵乘法(带转置)实现解析
【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym
transpose_quant_batch_matmul 是 CANN pypto-gym 仓库中基于 PyPTO 编程框架实现的实验性量化矩阵乘算子,它在一次 kernel 内完成「批量矩阵乘法 + MXFP8 块量化 + 多种 perm 转置组合 + FP16/BF16 输出」,并支持 M 轴动态化,是学习 PyPTO scaled_mm、E8M0 缩放因子布局与 tile 配置的典型样例。阅读本文后,你将掌握该算子的数学语义、ShapeConfig 各字段含义、permX2 与 b_trans 的映射关系、M 轴切分实现策略,以及如何通过仓库内测试用例与 Golden 实现完成精度验证。
算子定位与产品支持情况
该算子位于仓库 src/pypto_gym/ops/pypto_tensor/experimental/matmul/transpose_quant_batch_matmul/,与 gmm_mxfp8、quant_matmul_reduce_sum、quant_batch_matmul 等同属实验性 matmul 算子族(见 matmul 目录总览)。
产品支持情况如下:
| 产品 | 支持情况 |
|---|---|
| Ascend 950PR | 支持 |
| Atlas A3 训练系列 / Atlas A3 推理系列 | 不支持 |
| Atlas A2 训练系列 / Atlas A2 推理系列 | 不支持 |
注意:该算子仅面向 Ascend 950PR,测试用例也通过
@pytest.mark.soc("950")做了平台限定(见 test_transpose_quant_batch_matmul.py),在 A2/A3 平台上不可直接运行。
算子语义与数学公式
transpose_quant_batch_matmul实现基于 MXFP8 量化的批量矩阵乘法(带转置)。对每个 batch 索引 b(b ∈ [0, B-1]),计算公式为:
$$ out[:, b, :] = permY\left(permX1(x1)[:, b, :] \times permX2(x2)[b, :, :]\right) $$
其中各张量含义为:
x1形状[M, B, K],FP8 格式,M 轴动态;x2形状[B, K, N](permX2=[0,1,2])或[B, N, K](permX2=[0,2,1]);permX1=[1,0,2]:x1 从[M,B,K]变为[B,M,K];permX2=[0,1,2]或[0,2,1]:x2 的布局决定 scaled_mm 的 b_trans 参数;permY=[1,0,2]:输出从[B,M,N]变为[M,B,N]。
核心参数共 4 个:
| 参数名 | 说明 | 示例值 |
|---|---|---|
| M | 左矩阵的行维度(动态) | 8 ~ 32768 |
| K | 矩阵乘的公共维度 | 128 |
| N | 右矩阵的列维度 | 512 |
| B | batch 维度(编译期固定) | 128 |
- 输入 x1:形状
[M, B, K](FP8,M 轴动态); - 输入 x2:形状
[B, K, N]或[B, N, K](取决于 permX2); - 输出 out:形状
[M, B, N](FP16 或 BF16)。
输入输出规格
输入参数:
| 参数名 | 类型 | 形状 | 描述 |
|---|---|---|---|
| x1 | Tensor | [M, B, K] | 左矩阵(FP8E4M3 或 FP8E5M2) |
| x2 | Tensor | [B, K, N]或[B, N, K] | 右矩阵(FP8E4M3 或 FP8E5M2) |
| x1Scale | Tensor | [M, B, K//64, 2] | 左矩阵缩放因子(FP8E8M0) |
| x2Scale | Tensor | [B, K//64, N, 2]或[B, N, K//64, 2] | 右矩阵缩放因子(FP8E8M0) |
输出参数:
| 参数名 | 类型 | 形状 | 描述 |
|---|---|---|---|
| out | Tensor | [M, B, N] | 输出矩阵(FP16 或 BF16) |
MX 量化约束
⚠️ K 必须是 64 的倍数,这是 MXFP8 块量化格式的硬性要求。
调用前必须验证的约束检查清单:
- K 轴对齐:
K % 64 == 0; - perm 组合:permX2 决定 x2 和 x2Scale 的形状;
- Scale 形状:符合 MXFP8 格式要求;
- 输出 dtype:dtype=1→FP16,dtype=27→BF16;
- B 轴固定:batch_size 在编译期固定。
MXFP8 量化说明
MX 量化(Microscaling Quantization)是一种基于块缩放的量化格式:
- 量化块大小:每 64 个元素共享一个缩放因子;
- 缩放因子格式:FP8E8M0(8 位纯指数),隐含 mantissa=1.0;
- 数据格式:FP8E4M3FN 或 FP8E5M2;
- x1Scale:形状
[M, B, K//64, 2]; - x2Scale:形状取决于 permX2。
在仓库测试中,缩放因子的生成使用torch.float8_e8m0fnu类型,数据范围约束在[0.9, 1.1]附近,FP8 数据则通过torch.float8_e4m3fn/torch.float8_e5m2生成(见 test_transpose_quant_batch_matmul.py)。
perm 组合说明
permX2=[0, 1, 2](K,N 顺序)
- x2 形状:
[B, K, N]; - x2Scale 形状:
[B, K//64, N, 2]; - scaled_mm 参数:无 b_trans;
- 计算路径:
[M,K] × [K,N] → [M,N]。
permX2=[0, 2, 1](N,K 反序)
- x2 形状:
[B, N, K]; - x2Scale 形状:
[B, N, K//64, 2]; - scaled_mm 参数:
b_trans=True, scale_b_trans=True; - 计算路径:
[M,K] × [N,K]^T → [M,N]。
这一映射在 kernel 实现中直接体现:实现文件 transpose_quant_batch_matmul_impl.py 依据permX2 == [0,1,2]判断是否给pypto.scaled_mm传入b_trans=True, scale_b_trans=True,硬件内部完成转置,避免额外数据搬运。
核心实现:M 轴切分与非并行循环
kernel 入口为transpose_quant_batch_mat_mul_kernel,通过@pypto.frontend.jit装饰,其整体结构如下:
@pypto.frontend.jit( pass_options={ "auto_mix_partition": 1, "cube_l1_reuse_setting": {-1: 2}, "cube_nbuffer_setting": {-1: 4}, "vec_nbuffer_setting": {-2: 1, -1: 2}, }, runtime_options={"stitch_function_max_num": 1024, "device_sched_mode": 1}, ) def transpose_quant_batch_mat_mul_kernel(x1, x2, x1Scale, x2Scale, out, tile_config): ...ShapeConfig 数据结构
ShapeConfig是 dataclass,字段含义如下(见 transpose_quant_batch_matmul_impl.py):
| 字段 | 说明 | 默认值 |
|---|---|---|
| ori_shape | 原始形状[M, K, N],M 在 kernel 中动态 | 必填 |
| batch_size | B 轴大小(编译期固定) | 必填 |
| m_tile_shape | cube 运算 M 维 tile | 必填 |
| k_tile_shape | cube 运算 K 维 tile | 必填 |
| n_tile_shape | cube 运算 N 维 tile | 必填 |
| vector_tile_shape | 向量运算 tile 形状 | 必填 |
| num_k_groups | K 轴分组数 | 1 |
| num_n_groups | N 轴分组数 | 1 |
| in_dtype | 输入数据类型 | DT_FP8E4M3 |
| out_dtype | 输出数据类型 | DT_BF16 |
| permX1 | x1 的 perm | [1,0,2] |
| permX2 | x2 的 perm | [0,1,2] |
| permY | 输出 perm | [1,0,2] |
| description | 测试用例描述 | "" |
M 轴切分策略(避免 IR 爆炸)
实现采用「M 轴切分 + 非并行循环」策略,kernel docstring 中明确说明其动机:
- 外层 LOOP_M(非并行):将 M 切分为 m_chunk_size 的 tile 依次迭代;
- 内层 LOOP_B(非并行):遍历 B 个 batch;
- reshape:x1 重排为
[M, B*K],x1Scale 重排为[M, B*K//64, 2]; - 每次迭代从重排后的二维张量上按行区间
[m_begin:m_end]与列区间[begin:end]切片; - 将 mm_result 通过
pypto.assemble写入 local_out 的[m_begin, out_pos]; - 最终将 local_out
[M, B*N]reshape 回[M, B, N]。
选择非并行循环而非 parallel=True 的原因在于:并行循环会触发 LoopUnroll 与 ExpandFunction 展开,导致 IR 图爆炸、编译缓慢;而把循环迭代放到运行时处理,可以保持 IR 图小巧,加快编译速度。
核心循环体(batch 切分 + scaled_mm 调用)如下:
x1_reshape = pypto.reshape(x1, [M, B * K], inplace=True) x1_scale_reshape = pypto.reshape(x1Scale, [M, B * K // 64, 2], inplace=True) local_out = pypto.Tensor(shape=(M, B * N), dtype=out_dtype) for b_idx in range(B): begin = b_idx * K end = (b_idx + 1) * K x1_slice = x1_reshape[:, begin:end] x1_scale_slice = x1_scale_reshape[:, begin // 64:end // 64, :] x2_slice = x2[b_idx, :, :] x2_scale_slice = x2Scale[b_idx, :, :, :] if permX2 == [0, 1, 2]: mm_result = pypto.scaled_mm(x1_slice, x2_slice, out_dtype, x1_scale_slice, x2_scale_slice) else: mm_result = pypto.scaled_mm(x1_slice, x2_slice, out_dtype, x1_scale_slice, x2_scale_slice, b_trans=True, scale_b_trans=True) out_pos = b_idx * N pypto.assemble(mm_result, [0, out_pos], local_out) out[:, :, :] = pypto.reshape(local_out, [M, B, N])实现特点小结:
- MXFP8 量化:数据支持 FP8E4M3FN 和 FP8E5M2,缩放因子使用 FP8E8M0;
- M 轴动态化:同一编译 kernel 支持不同 M 大小,运行时通过
ori_shape[0]获取; - B 轴切分:非并行 LOOP_B 逐 batch 独立计算,避免 parallel 展开导致的 IR 膨胀;
- perm 组合支持:通过 permX2 决定 scaled_mm 的 b_trans 参数,硬件内完成转置;
- 输出 dtype 灵活:支持 FP16 和 BF16 输出;
- Cache 策略:输入 tensor 使用 NONE_CACHEABLE,减少 L1 缓存占用(MXFP8 数据一次性读取,无需缓存复用)。
调用示例
test1:小 M + b_trans 配置
test_transpose_quant_batch_matmul( ShapeConfig( ori_shape=[8, 128, 512], # [M, K, N] batch_size=128, # B轴大小 m_tile_shape=[256, 256], k_tile_shape=[128, 128], n_tile_shape=[256, 256], vector_tile_shape=[1, 128, 256, 32], in_dtype=pypto.DT_FP8E5M2, out_dtype=pypto.DT_BF16, permX1=[1, 0, 2], permX2=[0, 2, 1], # N,K反序 → b_trans=True permY=[1, 0, 2], description="test1" ) )test3:大 M + 标准配置
test_transpose_quant_batch_matmul( ShapeConfig( ori_shape=[8192, 128, 512], # [M, K, N] batch_size=128, m_tile_shape=[256, 256], k_tile_shape=[128, 128], n_tile_shape=[256, 256], vector_tile_shape=[1, 128, 256, 32], in_dtype=pypto.DT_FP8E4M3, out_dtype=pypto.DT_FP16, permX1=[1, 0, 2], permX2=[0, 1, 2], # K,N顺序 → 无b_trans permY=[1, 0, 2], description="test3" ) )注意:仓库测试文件 test_transpose_quant_batch_matmul.py 中实际注册的用例 tile 配置与 README 示例略有差异(如 test1 使用m_tile_shape=[128,128]、n_tile_shape=[512,512]),并且通过FILTERED_CONFIGS过滤掉了 test3(大 M 用例默认不在 pytest 参数化中执行,仅在__main__手动运行),从源码结构看这是为了控制常规 CI 测试时长。
优化配置参考
- NBuffer 配置:
cube_nbuffer_setting={-1:2},vec_nbuffer_setting={-2:1,-1:16}; - Cache 策略:
NONE_CACHEABLE用于所有输入 tensor; - Cube L1 复用:
cube_l1_reuse_setting={-1:2}。
(kernel 内实际 pass_options 为cube_nbuffer_setting={-1:4}、vec_nbuffer_setting={-2:1,-1:2},与 README 记录的历史优化配置略有出入,实际以 实现文件 为准。)
性能数据
| 配置 | M | K | N | B | 预估时间 |
|---|---|---|---|---|---|
| test1 | 8 | 128 | 512 | 128 | ~108 us |
| test2 | 128 | 128 | 512 | 128 | ~122 us |
| test3 | 8192 | 128 | 512 | 128 | ~9583 us |
| test4 | 32768 | 128 | 512 | 128 | ~40635 us |
以上为文档记录的预估时间(Ascend 950PR 平台),实际性能随环境、tile 配置与编译选项变化,应以实测为准。
精度验证与 Golden 实现
容差设置
- 相对容差(RTOL):1e-3;
- 绝对容差(ATOL):1e-3;
- 对比工具:
numpy.testing.assert_allclose。
验证方法
- Golden 实现:纯 PyTorch 实现,作为精度基准;
- 三态标记:
[PRECISION_PASS]或[PRECISION_FAIL]; - 对比工具:
numpy.testing.assert_allclose。
Golden 实现在 transpose_quant_batch_matmul_golden.py,其compute_golden_result完整复现了量化的逆过程,可作为理解 MXFP8 语义的参考实现:
- FP8 → FP32(
x.float()); - E8M0 Scale → FP32;
- 对 x1 应用
permute(permX1)([M,B,K] → [B,M,K]); - 对 x2 应用
permute(permX2); - 缩放因子 reshape 后按块用
torch.repeat_interleave(..., repeats=32)广播为逐元素缩放(每 64 元素一个缩放因子,展开为 2×32 结构); - 反量化:
x * scale; torch.matmul批量矩阵乘;- 输出
permute(permY)并按 dtype 标志转 FP16(dtype=1)或 BF16(dtype=27)。
测试用例(test_transpose_quant_batch_matmul.py)的整体流程为:生成 MXFP8 输入 → 计算 Golden → 包装器把张量搬到 NPU 并分配输出 → 调用 kernel →assert_allclose(golden, result, rtol=1e-3, atol=1e-3)比对。包装器transpose_quant_batch_matmul中 N 的取值逻辑也印证了 permX2 与形状的关系:N = x2.shape[-1] if permX2 == [0,1,2] else x2.shape[1]。
测试用例一览
| 测试名称 | M | K | N | B | permX2 | 说明 |
|---|---|---|---|---|---|---|
| test1 | 8 | 128 | 512 | 128 | [0,2,1] | 小 M+b_trans 验证 |
| test2 | 128 | 128 | 512 | 128 | [0,1,2] | 中 M+标准验证 |
| test3 | 8192 | 128 | 512 | 128 | [0,1,2] | 大 M 性能验证 |
| test4 | 32768 | 128 | 512 | 128 | [0,1,2] | 超大规模验证 |
常见问题
Q1:为什么 permX2=[0,2,1] 需要 b_trans=True?
- permX2=[0,2,1] 表示 x2 的布局是
[B,N,K]; b_trans=True让 scaled_mm 在硬件内部完成转置;scale_b_trans=True同时处理 scale tensor 的转置;- 这样避免了额外的数据搬运操作。
Q2:M 轴动态化如何实现?
- 编译期不固定 M 值,同一 kernel 可处理不同 M 的输入;
- B 轴在编译期固定(batch_size 参数),通过循环切分;
- M 轴动态化减少重编译开销,适配不同输入规模;
- 从源码实现看,M 通过
tile_config.ori_shape[0]在运行时读取,配合非并行 LOOP 切片完成。
Q3:为什么输入 tensor 使用 NONE_CACHEABLE?
- MXFP8 数据一次性读取,无需缓存复用;
- NONE_CACHEABLE 减少 L1 缓存占用,为输出 tensor 腾出空间;
- 对单次读取场景无性能损失。
参考文档
- SPEC.md- 详细需求规格
- API_REPORT.md- API 映射分析
- DESIGN.md- 详细设计文档
- matmul 算子族总览
- kernel 实现
- 精度验证测试
- Golden 参考实现
版本历史
| 版本 | 日期 | 说明 |
|---|---|---|
| v1.0 | 2026-06-03 | 初始版本,支持 MXFP8 批量矩阵乘法(带转置) |
【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考