CANN pypto-gym 实战:transpose_quant_batch_matmul 算子的 MXFP8 量化批量矩阵乘法(带转置)实现解析
2026/9/18 19:58:32 网站建设 项目流程

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
Bbatch 维度(编译期固定)128
  • 输入 x1:形状[M, B, K](FP8,M 轴动态);
  • 输入 x2:形状[B, K, N][B, N, K](取决于 permX2);
  • 输出 out:形状[M, B, N](FP16 或 BF16)。

输入输出规格

输入参数:

参数名类型形状描述
x1Tensor[M, B, K]左矩阵(FP8E4M3 或 FP8E5M2)
x2Tensor[B, K, N][B, N, K]右矩阵(FP8E4M3 或 FP8E5M2)
x1ScaleTensor[M, B, K//64, 2]左矩阵缩放因子(FP8E8M0)
x2ScaleTensor[B, K//64, N, 2][B, N, K//64, 2]右矩阵缩放因子(FP8E8M0)

输出参数:

参数名类型形状描述
outTensor[M, B, N]输出矩阵(FP16 或 BF16)

MX 量化约束

⚠️ K 必须是 64 的倍数,这是 MXFP8 块量化格式的硬性要求。

调用前必须验证的约束检查清单:

  1. K 轴对齐K % 64 == 0
  2. perm 组合:permX2 决定 x2 和 x2Scale 的形状;
  3. Scale 形状:符合 MXFP8 格式要求;
  4. 输出 dtype:dtype=1→FP16,dtype=27→BF16;
  5. 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_sizeB 轴大小(编译期固定)必填
m_tile_shapecube 运算 M 维 tile必填
k_tile_shapecube 运算 K 维 tile必填
n_tile_shapecube 运算 N 维 tile必填
vector_tile_shape向量运算 tile 形状必填
num_k_groupsK 轴分组数1
num_n_groupsN 轴分组数1
in_dtype输入数据类型DT_FP8E4M3
out_dtype输出数据类型DT_BF16
permX1x1 的 perm[1,0,2]
permX2x2 的 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])

实现特点小结:

  1. MXFP8 量化:数据支持 FP8E4M3FN 和 FP8E5M2,缩放因子使用 FP8E8M0;
  2. M 轴动态化:同一编译 kernel 支持不同 M 大小,运行时通过ori_shape[0]获取;
  3. B 轴切分:非并行 LOOP_B 逐 batch 独立计算,避免 parallel 展开导致的 IR 膨胀;
  4. perm 组合支持:通过 permX2 决定 scaled_mm 的 b_trans 参数,硬件内完成转置;
  5. 输出 dtype 灵活:支持 FP16 和 BF16 输出;
  6. 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 记录的历史优化配置略有出入,实际以 实现文件 为准。)

性能数据

配置MKNB预估时间
test18128512128~108 us
test2128128512128~122 us
test38192128512128~9583 us
test432768128512128~40635 us

以上为文档记录的预估时间(Ascend 950PR 平台),实际性能随环境、tile 配置与编译选项变化,应以实测为准。

精度验证与 Golden 实现

容差设置

  • 相对容差(RTOL):1e-3;
  • 绝对容差(ATOL):1e-3;
  • 对比工具:numpy.testing.assert_allclose

验证方法

  1. Golden 实现:纯 PyTorch 实现,作为精度基准;
  2. 三态标记[PRECISION_PASS][PRECISION_FAIL]
  3. 对比工具numpy.testing.assert_allclose

Golden 实现在 transpose_quant_batch_matmul_golden.py,其compute_golden_result完整复现了量化的逆过程,可作为理解 MXFP8 语义的参考实现:

  1. FP8 → FP32(x.float());
  2. E8M0 Scale → FP32;
  3. 对 x1 应用permute(permX1)[M,B,K] → [B,M,K]);
  4. 对 x2 应用permute(permX2)
  5. 缩放因子 reshape 后按块用torch.repeat_interleave(..., repeats=32)广播为逐元素缩放(每 64 元素一个缩放因子,展开为 2×32 结构);
  6. 反量化:x * scale
  7. torch.matmul批量矩阵乘;
  8. 输出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]

测试用例一览

测试名称MKNBpermX2说明
test18128512128[0,2,1]小 M+b_trans 验证
test2128128512128[0,1,2]中 M+标准验证
test38192128512128[0,1,2]大 M 性能验证
test432768128512128[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.02026-06-03初始版本,支持 MXFP8 批量矩阵乘法(带转置)

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

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

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

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

立即咨询