CANN 昇腾 MXFP8 量化矩阵乘算子实战:SWAT 模板、K-Tail Stepwise Copyout 与自动化调优选型
CANN 昇腾 MXFP8 量化矩阵乘算子实战:SWAT 模板、K-Tail Stepwise Copyout 与自动化调优选型
MXFP8(Microscaling FP8)是一种 8 位浮点数量化格式,通过共享缩放因子(scale)在保持较好精度的前提下显著降低带宽开销与访存压力,尤其适合大语言模型推理等带宽敏感场景。本文以 cann-samples 仓库中 quant_matmul_mxfp8 示例 为骨架,完整讲解其在昇腾 AI 处理器(NPU ARCH 3510)上的实现、命令行参数、一键/手动运行流程,并结合源码剖析 SWAT 滑动窗口模板、K-Tail Stepwise Copyout、4-buffer、A 全载与权重 NZ 排布等优化方案的底层实现与算法自动推荐机制。读完本文,你将掌握如何在昇腾 NPU 上编译、运行、校验并自动选型 MXFP8 量化矩阵乘算子,以及如何利用仓库自带的 msprof 链路对多个候选算法做兼容性筛选与耗时排序。
一、样例能力全景
该示例位于 matmul_recipes/examples/quant_matmul_mxfp8,围绕 MXFP8 量化矩阵乘提供了 5 个可执行目标与 5 个辅助脚本,覆盖权重 ND / NZ 两种数据排布:
可执行目标(权重 ND 排布,含 NDExtLayout / DNExtLayout)
| 可执行目标 | 说明 |
|---|---|
quant_matmul_mxfp8_swat |
基于 SWAT 模板、双 L1 缓冲(2-buffer)的实现 |
quant_matmul_mxfp8_k_tail_stepwise_copyout |
基于 SWAT non-full-load 路径的 K-Tail Stepwise Copyout 实现 |
quant_matmul_mxfp8_swat_4_buffer |
基于 SWAT 模板、四 L1 缓冲(4-buffer)的实现 |
quant_matmul_mxfp8_a_full_load |
A 矩阵 full load 方案的实现 |
可执行目标(权重 NZ 排布,B 为 NZLayout / ZNLayout,A 排布与 ND 一致)
| 可执行目标 | 说明 |
|---|---|
quant_matmul_mxfp8_swat_weight_nz |
基于 SWAT 模板、双 L1 缓冲(2-buffer)的实现 |
辅助脚本(scripts/ 目录)
| 脚本 | 作用 |
|---|---|
gen_data.py |
生成 ND 权重输入数据和 CPU golden 结果 |
gen_data_weight_nz.py |
生成 NZ 权重输入数据和 CPU golden 结果 |
verify_result.py |
校验 NPU 输出与 CPU golden 是否一致 |
quant_matmul_mxfp8_algorithm_recommend.py |
对当前目录下可执行算法进行兼容性筛选和耗时排序 |
run.sh |
一键串联构建、数据生成、算子执行与校验 |
每个可执行目标本质上是一个独立的 launcher(.asc 源文件),它们共享同一套 host tiling 与 kernel 模板实现,仅在调度策略、缓冲深度与搬运方式上分化。以 quant_matmul_mxfp8_swat.asc 为例,其核心类型与策略组合为:
using TypeA = fp8_e4m3fn_t; // A 矩阵为 FP8 E4M3
using TypeB = fp8_e4m3fn_t; // B 矩阵为 FP8 E4M3
using TypeC = bfloat16_t; // 输出累加结果落盘为 BF16
using TypeScaleA = fp8_e8m0_t; // A 的量化缩放因子为 FP8 E8M0
using TypeScaleB = fp8_e8m0_t;
constexpr int32_t SCALE_C0 = 2; // scale 在 C0 维的宽度为 2
using BlockScheduler = QuantMatmulMxSwatScheduler<NO_FULL_LOAD_MODE>;
using DispatchPolicy = QuantMatmulMxMultiBlockWithSwat<NO_FULL_LOAD_MODE, 2UL>;
从源码结构看,MXFP8 量化矩阵乘采用了"数据 + 缩放因子"双输入范式:A/B 矩阵以 fp8_e4m3fn_t 存储,scaleA/scaleB 以 fp8_e8m0_t 存储,SCALE_C0=2 与 MX 格式每 32 个元素共享 2 个 scale 的布局相对应;输出 C 使用 bfloat16_t,在精度与带宽之间取得平衡。
二、使用约束与支持架构
当前样例支持以下场景:
- 支持通过命令行参数
transA/transB选择 A/B 矩阵转置; - 支持权重矩阵 B 数据排布
ND/NZ输入; - 支持 NPU 架构:NPU ARCH 3510(构建时对应
-DNPU_ARCH=dav-3510)。
其中 NZ 排布仅由 quant_matmul_mxfp8_swat_weight_nz 支持,且其 A 矩阵数据排布与 ND 系列保持一致。
三、命令行参数说明
所有可执行文件的命令行参数格式一致:
<program> m k n [transA transB]
| 参数 | 说明 |
|---|---|
m |
矩阵 A 的行数 |
k |
矩阵 A 的列数,同时也是矩阵 B 的归约维 |
n |
矩阵 B 的行数,对应输出矩阵的列数 |
transA(可选) |
A 矩阵转置信息(0/1/true/false/t/f)。0/false/f 表示非转置,shape 为 [M, K];1/true/t 表示转置,shape 为 [K, M]。默认为非转置。 |
transB(可选) |
B 矩阵转置信息(0/1/true/false/t/f)。0/false/f 表示非转置,shape 为 [K, N](ND)/ [N1, K1, K0, N0](NZ);1/true/t 表示转置,shape 为 [N, K](ND)/ [K1, N1, N0, K0](NZ)。其中 K0=32、N0=16,K1=ceil(K/K0),N1=ceil(N/N0)。默认为转置。 |
输出矩阵 C 的逻辑形状为 [M, N]。
在 run.sh 与 quant_matmul_mxfp8_algorithm_recommend.py 的源码中,转置标志的解析逻辑一致:0/false/f 归一化为 0,1/true/t 归一化为 1,其余输入直接报错并提示合法取值;m k n 为必填参数,transA/transB 要么同时省略(默认 transA=false、transB=true),要么同时给出。
四、数据生成与结果校验
4.1 输入输出文件
gen_data.py 与 gen_data_weight_nz.py 会在运行时的当前工作目录(通常为安装目录 build_out/.../quant_matmul_mxfp8)下生成以下文件:
input/input_a.bin
input/input_b.bin
input/input_scaleA.bin
input/input_scaleB.bin
output/cpu_output.bin
注意:两个数据生成脚本输出的文件名完全相同,仅 B 矩阵数据排布方式不同(ND vs NZ),请勿混用与可执行目标不匹配的数据。
样例执行完成后会额外生成 output/npu_out.bin。每个可执行文件在运行结束后都会自动调用 verify_result.py,将 NPU 输出与 CPU golden 进行一致性校验(launcher 源码末尾通过 python3 verify_result.py m n 触发)。
4.2 MXFP8 量化与 golden 生成逻辑
从 gen_data.py 源码可以看到 MXFP8 数据的具体形态:
- A/B 元素以
float8_e4m3fn(ml_dtypes)生成,取值区间[1, 8); - scale 以
float8_e8m0(en_dtypes)生成,shape 遵循 MX 分组规则:transA=False时a_scale=[m, ceil(k/64), 2],transA=True时为[ceil(k/64), m, 2];B 侧同理,transB=True时b_scale=[n, ceil(k/64), 2]; - golden 计算流程为:先按 scale 反量化(
dequant_mxfp8,将 scale 沿最后一维 repeat 32 次并沿ceil(x/64)维展开),再按转置属性交换轴后执行torch.matmul,最后转成torch.bfloat16落盘; - MXFP8 每个数值占 1 字节,A/B 直接以
np.uint8视角写盘,输出则以torch.uint16位模式写盘(BF16 的 16 位原始载荷)。
gen_data_weight_nz.py 额外实现了 to_weight_nz_layout:将 2D [row, col] 的 B 矩阵转换为 GM 侧 weight-NZ 布局 [ceil(col/32), ceil(row/16), 16, 32](FP8_C0_SIZE=32、CUBE_BLOCK=16),即在 32×16 的小块内按 (col_tile, row_tile, 16, 32) 重排,与 README 中 NZ shape [K1, N1, N0, K0] / [N1, K1, K0, N0] 的定义保持一致。
4.3 校验容差
verify_result.py 采用相对误差与绝对误差双重判据:
POINT_ERROR_TOL = 1e-1(单点相对误差阈值);RATIO_POINT_ERROR_TOL = 1e-3(绝对误差阈值);ERROR_RATIO_TOL = 1e-3(整体错误比例阈值)。
对于 m*n 较大的矩阵,脚本会跳过全量打印,改为输出 abs_err / rel_err / RMSE 的统计摘要与左上角 4×4 角块对比,便于快速定位误差分布。
五、一键运行(推荐)
仓库提供 run.sh,可一键串联构建 → 数据生成 → 算子执行 → 结果校验全流程。推荐先进入样例目录再执行:
cd Samples/2_Performance/matmul_story/matmul_recipes/examples/quant_matmul_mxfp8
# 自动构建 + 自动推荐最优算法 + 运行
bash scripts/run.sh 16 128 16384 0 1
# 指定 K-Tail Stepwise Copyout 目标,跳过重新构建
bash scripts/run.sh \
--target quant_matmul_mxfp8_k_tail_stepwise_copyout --skip-build 4096 1024 32768 0 1
# 查看完整帮助
bash scripts/run.sh --help
run.sh 参数说明
| 参数 | 说明 |
|---|---|
m k n [transA transB] |
矩阵维度与转置参数。transA/transB 可选,支持 0/1/true/false/t/f;省略时默认 transA=false(0)、transB=true(1)。 |
--target <name> |
指定要运行的可执行文件名。省略时自动调用推荐脚本选择最优目标。 |
--skip-build |
跳过构建/安装阶段,复用已有 build_out。 |
-h, --help |
显示帮助信息。 |
从 run.sh 源码看,其内部执行顺序为:向上查找仓库根(以 .ci/build.sh 为锚点)→ 默认调用 .ci/build.sh 构建安装 → 进入安装目录 build_out/2_Performance/matmul_story/matmul_recipes/quant_matmul_mxfp8 → 未指定 --target 时调用 quant_matmul_mxfp8_algorithm_recommend.py m k n transA transB --print-target 自动选择目标 → 依据目标名是否以 _weight_nz 结尾选择 gen_data_weight_nz.py 或 gen_data.py 生成数据 → 运行可执行文件(其中已包含自动校验)。
如需查看完整算法推荐排名(含耗时表格),请在安装目录下直接运行 quant_matmul_mxfp8_algorithm_recommend.py(见下文"手动构建与运行")。
六、手动构建与运行
如需手动控制各步骤,可在仓库根目录下完成编译和安装后,进入当前样例目录:
cmake -S . -B build -DNPU_ARCH=dav-3510
cmake --build build --parallel
cmake --install build --prefix ./build_out
cd build_out/2_Performance/matmul_story/matmul_recipes/quant_matmul_mxfp8
6.1 生成测试数据
ND 权重:
python3 gen_data.py 16 128 16384 0 1
NZ 权重:
python3 gen_data_weight_nz.py 16 128 16384 0 1
6.2 运行单个算法样例
ND 权重:
# 2-buffer SWAT
./quant_matmul_mxfp8_swat 16 128 16384 0 1
# 或:4-buffer SWAT
./quant_matmul_mxfp8_swat_4_buffer 16 128 16384 0 1
# 或:A full load
./quant_matmul_mxfp8_a_full_load 16 128 16384 0 1
# K-Tail Stepwise Copyout:重新生成匹配该 Shape 的输入数据
python3 gen_data.py 4096 1024 32768 0 1
./quant_matmul_mxfp8_k_tail_stepwise_copyout 4096 1024 32768 0 1
NZ 权重:
# 2-buffer SWAT
./quant_matmul_mxfp8_swat_weight_nz 16 128 16384 0 1
每个可执行文件运行结束后都会自动调用 verify_result.py 完成 NPU 输出与 CPU golden 的一致性校验,退出码非 0 即表示校验失败。
6.3 运行算法推荐脚本
python3 quant_matmul_mxfp8_algorithm_recommend.py 16 128 16384 0 1
下图为推荐脚本输出的结构示意(数值为虚构,仅说明版式):
[Profile Breakdown]
+------------------------------------------------+----------+---------+----------+---------+---------+------------+--------------+
| candidate |kernel(us)| mac(us) |scalar(us)| mte1(us)| mte2(us)|fixpipe(us) |icache_miss(%)|
+================================================+==========+=========+==========+=========+=========+============+==============+
| quant_matmul_mxfp8_swat_weight_nz | 10.100| 1.150 | 0.520| 0.100 | 0.280 | 0.720 | 0.082 |
| quant_matmul_mxfp8_k_tail_stepwise_copyout | 10.500| 1.170 | 0.530| 0.105 | 0.320 | 0.730 | 0.086 |
| quant_matmul_mxfp8_swat_4_buffer | 11.900| 1.200 | 0.550| 0.110 | 0.440 | 0.770 | 0.095 |
| quant_matmul_mxfp8_swat | 12.345| 1.234 | 0.567| 0.123 | 0.456 | 0.789 | 0.100 |
| quant_matmul_mxfp8_a_full_load | 15.678| 2.100 | 0.800| 0.200 | 0.300 | 0.500 | 0.250 |
+------------------------------------------------+----------+---------+----------+---------+---------+------------+--------------+
七、算法自动推荐机制:从 msprof 到排序
quant_matmul_mxfp8_algorithm_recommend.py 是仓库提供的一个可复用的调优选型工具,其核心链路为:
- 候选发现(自动扩展):扫描安装目录下所有可执行文件(POSIX 无后缀可执行或
.exe),将每个可执行目标视为一个候选算法,排除脚本自身。这样新增算法时无需修改推荐脚本本身。 - 候选分组:以
_weight_nz后缀为标识,将候选划分为 ND 权重与 NZ 权重两组;ND 组用gen_data.py生成输入,NZ 组用gen_data_weight_nz.py生成输入,保证同组候选在完全相同的输入数据与 Shape 条件下比较。 - msprof 逐候选剖析:对每个候选执行
msprof --output=... <program> m k n transA transB,随后在PROF_*目录下的mindstudio_profiler_output/op_summary_*.csv中提取指标。 - 指标提取:读取
Task Duration(us)作为 kernel 总耗时,并同步提取aic_mac_time(us)、aic_scalar_time(us)、aic_mte1_time(us)、aic_mte2_time(us)、aic_fixpipe_time(us)、aic_icache_miss_rate等细分指标(对应表格中的 kernel/mac/scalar/mte1/mte2/fixpipe/icache_miss 列)。 - 兼容性筛选与排序:运行失败(非零退出码)或无法解析出完整指标的候选被过滤掉;其余候选按 kernel 耗时升序排列,输出
[Recommended Algorithm Ranking]与[Profile Breakdown]表格。 --print-target模式:只输出最优候选的可执行文件名(推荐信息打印到 stderr),供run.sh直接捕获并调用。
值得注意的一个实现细节:脚本通过 PROFILING_KERNEL_RUN_INDEX = 1 定位 op_summary 中第二次 kernel 启动的数据行——这是因为各示例 launcher 使用 EXAMPLE_KERNEL_RUN_COUNT=2(一次 warmup + 一次正式测量),取第二次测量值可避免冷启动与 cache 预热带来的抖动。
八、各优化实现的技术要点
8.1 SWAT 模板(quant_matmul_mxfp8_swat / quant_matmul_mxfp8_swat_weight_nz)
SWAT(Slide Window Adaptive Template,自适应滑动窗口模板)通过提升多核单次访问的 L2 命中率来提高 MTE2 搬运效率,从而实现首轮搬运即可做到 MMAD 指令不断流,使算子在 Cube Bound 场景下计算单元利用率达到 95%+(详见 MX 量化矩阵乘算子性能优化指南 中"SWAT"一节)。源码中对应:
using BlockScheduler = QuantMatmulMxSwatScheduler<NO_FULL_LOAD_MODE>;
using DispatchPolicy = QuantMatmulMxMultiBlockWithSwat<NO_FULL_LOAD_MODE, 2UL>;
NO_FULL_LOAD_MODE 表示非全载路径,2UL 为 L1 双缓冲深度。quant_matmul_mxfp8_swat_weight_nz 与其差异仅在于权重 B 走 NZLayout/ZNLayout,A 与 scale 的排布保持一致。
8.2 K-Tail Stepwise Copyout(quant_matmul_mxfp8_k_tail_stepwise_copyout)
当完整输出基本块大于半个 L0C 时,该实现保持 M 方向不变,仅在首个和最后一个 KL1 轮次沿 N 方向切分;K 尾轮中每个 N 子块完成全部 KL0 累加后立即搬出,以减少 FIXPIPE 读取与下一 Tile MMAD 写入之间的冲突。其 DispatchPolicy 为:
using DispatchPolicy = QuantMatmulMxKTailStepwiseCopyout<NO_FULL_LOAD_MODE, 2UL>;
其动机在 性能优化指南 中有详细阐述:UnitFlag 可以让 FIXPIPE 在当前 Tile 最终累加期间提前搬出已完成的部分结果,但当 FIXPIPE 的搬出尾部延续到下一 Tile 时,下一 Tile 立即复用相同 L0C 地址仍可能产生读写冲突。K-Tail Stepwise Copyout 将 K 尾轮循环顺序从 KL1 -> KL0 -> N Split 调整为 KL1 -> N Split -> KL0:第一个 N 子块先连续完成全部 KL0 累加并在最终累加完成后立即触发搬出,再执行第二个 N 子块的整段 KL0 计算,使同一 L0C 子区域在被下一 Tile 复用前获得更长的计算间隔,从而减少冲突并提高 K 尾轮计算与结果搬出的并行度。README 示例中该目标使用更大 Shape(如 4096 1024 32768)运行,因为其针对大 K / 多计算轮次、FIXPIPE 流水空洞在连续 Tile 间累积的场景。
8.3 4-Buffer(quant_matmul_mxfp8_swat_4_buffer)
在双缓冲基础上将 L1 缓冲深度提升到 4,对应 kernel 实现 Kernel::QuantMatmulMxKernelSwat4Buffer。从性能指南看,使能 4-Buffer 的条件是:Profiling 显示存在明显的流水停顿(计算等待搬运),且 L1/L0 空间评估后仍能保持目标 base 块大小、不引入新的 MTE1/FIXPIPE 瓶颈。
8.4 A Full Load(quant_matmul_mxfp8_a_full_load)
对应 kernel 实现 Kernel::QuantMatmulMxKernelAFullLoad,在 MTE2 Bound 场景中通过减少 MTE2 的整体搬运量来提升性能:当输入可完整缓存在 L1 中时,全载模板使输入始终驻留 L1,从而减少整体搬运耗时。当前样例实现仅支持 A 全载(左矩阵 a 及其 scaleA 常驻 L1),暂不支持 B 全载。适用场景为:MTE2 Bound 为主要瓶颈、输入矩阵较小可完全载入 L1、Decode 等小批量计算场景。
8.5 与 MX 量化性能指南的衔接
上述各方案的原理、指令级细节(如 UnitFlag 的 512B 粒度同步、L1 Bank 冲突优化、Scale 缓存、尾轮负载均衡、性能建模公式与 Bound 判定)均在 MX 量化矩阵乘算子性能优化指南 中有系统性展开,该文档也是本示例 README 推荐的首选延伸阅读,包含:
- 算子实现原理与 Tensor a/b/scaleA/scaleB 的搬运说明(GM→L1 的布局变换、
ND2NZ/DN2NZ指令与 K 方向补零约束); - 性能建模:FIXPIPE vs MMAD 的 Bound 判定、流水理论耗时对比;
- 优化实践:Double Buffer、UnitFlag、K-Tail Stepwise Copyout、SWAT、L1 Bank 冲突、Scale 缓存、全载优化、尾轮负载均衡;
- 优化策略选择指南:按 Bound 类型(Cube Bound / FIXPIPE Bound / MTE2 Bound)与输入特征(大规模矩阵、小批量 Decode 等)选择模板与优化手段。
九、总结
quant_matmul_mxfp8 示例以极小的学习成本展示了 MXFP8 量化矩阵乘在昇腾 NPU 上的完整落地链路:从数据生成(含 MX 量化、scale 布局、ND/NZ 排布)到多种模板/策略的实现分化(SWAT 2-buffer / 4-buffer、K-Tail Stepwise Copyout、A full load、weight-NZ),再到基于 msprof 的自动化选型。读者既可以通过 run.sh 一键体验全流程,也可以手动构建运行逐一对比各算法的 mac/scalar/mte1/mte2/fixpipe/icache_miss 明细指标,并结合 MX 量化矩阵乘算子性能优化指南 深入理解每个优化手段背后的硬件流水原理,从而在实际业务中按 Shape 特征与瓶颈类型做出正确的模板与策略选择。