CANN 昇腾 MXFP8 量化矩阵乘算子实战:SWAT 模板、K-Tail Stepwise Copyout 与自动化调优选型

原创2026-09-17 23:53:221,151 阅读
文章标签:示例工程CANN

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 是仓库提供的一个可复用的调优选型工具,其核心链路为:

  1. 候选发现(自动扩展):扫描安装目录下所有可执行文件(POSIX 无后缀可执行或 .exe),将每个可执行目标视为一个候选算法,排除脚本自身。这样新增算法时无需修改推荐脚本本身。
  2. 候选分组:以 _weight_nz 后缀为标识,将候选划分为 ND 权重与 NZ 权重两组;ND 组用 gen_data.py 生成输入,NZ 组用 gen_data_weight_nz.py 生成输入,保证同组候选在完全相同的输入数据与 Shape 条件下比较。
  3. msprof 逐候选剖析:对每个候选执行 msprof --output=... <program> m k n transA transB,随后在 PROF_* 目录下的 mindstudio_profiler_output/op_summary_*.csv 中提取指标。
  4. 指标提取:读取 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 列)。
  5. 兼容性筛选与排序:运行失败(非零退出码)或无法解析出完整指标的候选被过滤掉;其余候选按 kernel 耗时升序排列,输出 [Recommended Algorithm Ranking] 与 [Profile Breakdown] 表格。
  6. --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 特征与瓶颈类型做出正确的模板与策略选择。

登录后查看全文
cann-samples