量化稀疏MLA算子测试框架实战指南:quant_sparse_flash_mla 的 pytest 精度验证体系 量化稀疏MLA算子测试框架实战指南quant_sparse_flash_mla 的 pytest 精度验证体系【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerquant_sparse_flash_mla 是 CANN ops-transformer 仓库中用于稀疏闪存 MLAMulti-head Latent Attention量化的 NPU 算子本文围绕其测试框架展开系统讲解基于 pytest 的功能验证方案CPU 侧 golden 复现、NPU 侧算子直调与图模式调用、以及 CPU 与 NPU 结果的精度对比覆盖单用例调测与 Excel 驱动批量测试两种模式。读完本文你将掌握该算子的完整测试流程、全部参数语义与命令选项并能独立编写与执行自己的测试用例。测试框架总体设计量化稀疏 MLA 算子的验证难点在于输入涉及量化反量化q/descale、稀疏索引sparse indices、两种 KV原始 ori_kv 与压缩 cmp_kv、多种 mask 模式以及可选的 softmax LSE 输出结果对精度极其敏感。为此tests/pytest 目录下实现了一套三阶段验证链路CPU 侧用 PyTorch 算子复现quant_sparse_flash_mla的完整计算流程生成 golden 数据真值实现位于 quant_sparse_flash_mla_golden.pyNPU 侧通过torch_npu完成算子直调eager 模式torch.ops.cann_ops_transformer.quant_sparse_flash_mla以及 aclgraph 图模式调用torch.compile torchair获取实际输出调用封装见 batch/quant_sparse_flash_mla_process.py精度对比CPU golden 与 NPU 输出统一送入 result_compare_method.py 进行逐元素对比输出通过率与详细差异报告。三阶段由 pytest 统一编排测试标记在 pytest.ini 中声明并通过 test_run.sh 一键执行。当前实现范围与参数限制框架目前已覆盖的算子形态如下与 check_valid_param.py 中的合法参数校验逻辑一一对应维度支持范围说明layout_qBSND、TNDQ 张量排布校验代码见 check_valid_param.pylayout_kvBSND、TND、PA_BBNDKV 张量排布其中PA_BBND为分页块排布q_typetorch.uint8Q 输入仅支持 uint8 量化格式hif8template_run_modeCSA、HCA、SWA、ORI_SPARSE、ORI_CMP_SPARSE稀疏注意力模板模式需要特别说明的是当layout_q为TND时cu_seqlens_q与T1均不能为None且cu_seqlens_q[-1]必须等于T1长度须等于B 1当layout_kv为TND时cu_seqlens_ori_kv必须单调递增cu_seqlens_cmp_kv若给定同样要求单调递增ORI_SPARSE/ORI_CMP_SPARSE模板下ori_mask_mode仅支持 0/3/4cmp_mask_mode仅支持 0/3。这些约束均在实际运行前由check_valid_param统一抛出ValueError被单用例主程序捕获后以pytest.skip跳过非法组合避免无效用例进入执行阶段。环境配置前置要求确认torch_npu为最新版本测试脚本通过torch_npu.npu.set_device指定 NPU 设备并依赖其完成张量的.npu()迁移与同步编译并安装本仓库算子即cann_ops_transformerbatch 进程通过import cann_ops_transformer注册torch.ops.cann_ops_transformer.quant_sparse_flash_mla与quant_sparse_flash_mla_metadata两个自定义算子入口。custom 包调用框架支持通过 custom 包调用算子即算子以自定义包形式注册进torch.ops命名空间后被测试代码直接使用单用例与批量测试统一走该入口。文件结构说明测试目录 tests/pytest 的文件组织如下tests/pytest/ ├── test_run.sh # 用例执行脚本single/batch_save/batch_exec/batch ├── quant_sparse_flash_mla_golden.py # 算子入参处理及 CPU 侧 golden 实现 ├── result_compare_method.py # CPU golden 与 NPU 结果精度对比脚本 ├── pytest.ini # 创建测试标记ci/graph/consistency ├── check_valid_param.py # 输入参数合法性校验 ├── generate_hifloat8_data.py # hif8 量化数据生成与反量化工具 ├── utils.py # 参数填充、结果保存等公共工具 ├── test_quant_sparse_flash_mla_single.py # pytest 单用例运行主程序 ├── quant_sparse_flash_mla_paramset.py # 单用例入参数配置 ├── test_quant_sparse_flash_mla_batch.py # 用例批量测试主程序读取 pt 文件执行 NPU 测试 └── batch/ ├── test_quant_sparse_flash_mla_pt_save.py # 读取 excel 表格批量生成用例 pt 文件 └── quant_sparse_flash_mla_process.py # 调用算子接口获取 NPU 输出其中single模式使用quant_sparse_flash_mla_paramset.py中手工配置的参数集批量模式则以 Excel 表格为用例来源先由test_quant_sparse_flash_mla_pt_save.py完成 golden 计算并落盘为.pt文件再由test_quant_sparse_flash_mla_batch.py加载 pt 文件执行 NPU 测试结果保存至 Excel 文件。单用例测试配置参数集手动编辑 quant_sparse_flash_mla_paramset.py其中每个键对应一组用例参数每个值是一个列表框架会对其做笛卡尔积展开形成多组参数组合最后通过ENABLED_PARAMS决定启用哪些参数集ENABLED_PARAMS [TEST_PARAMS[decode_first]] # ENABLED_PARAMS [TEST_PARAMS[key] for key in TEST_PARAMS.keys()]内置参数集包括decode_first解码首 token、prefill_first预填充、csa_small_prefill小规模预填充、ori_sparse_tnd_pa/ori_sparse_bsnd_pa仅原始稀疏、ori_cmp_sparse_tnd_pa/ori_cmp_sparse_bsnd_pa原始压缩双稀疏、load_4096_swa_tnd_pa滑动窗口 SWA、load_8192_hca_tnd_tndHCA 模板等覆盖不同 layout、模板模式与 batch 一致性场景。核心参数语义如下参数语义典型取值layout_q/layout_kvQ / KV 张量排布BSND、TND、PA_BBNDq_type/ori_kv_type/cmp_kv_type量化数据类型torch.uint8B/S1/S2batch、Q 序列长度、KV 序列长度如B1, S11, S28192N1/N2/D/KQ 头数、KV 头数、头维、topk 数如N164, N21, D512, K512block_num1/block_num2分页块数PA 布局None自动推导block_size1/block_size2分页块大小128seqused_q/seqused_ori_kv/seqused_cmp_kv各 batch 实际使用长度如[2, 2]、[4096]cu_seqlens_q/cu_seqlens_ori_kv/cu_seqlens_cmp_kv前缀和序列长度TND 布局必须None或前缀和列表cmp_residual_kv压缩 KV 的残余长度如[0]、[2]、[3]softmax_scalesoftmax 缩放系数0.04419417即 1/√512cmp_ratio压缩比ori_kv 与 cmp_kv 长度比4或1ori_mask_mode/cmp_mask_modemask 模式0无 mask、3causal、4滑动窗口0/3/4与0/3ori_win_left/ori_win_right滑动窗口左右边界mask 模式 4-1表示不限制quant_mode量化模式1template_run_mode稀疏模板CSA/HCA/SWA/ORI_SPARSE/ORI_CMP_SPARSEactlen_mode实际长度模式fullS1EQS2S1 是否等于 S2FalseisSink是否启用 sink token 优化Truereturn_softmax_lse是否返回 softmax LSETrue/Falseori_kv_topk_mode/cmp_kv_topk_modetopk 生成方式fullK/random/noori_sparse_indices_mode/cmp_sparse_indices_mode稀疏索引生成方式full/randomori_topk_length/cmp_topk_lengthtopk 长度None自动计算batch_consistency*确定性三级 batch 一致性校验相关见下文注意Testcase_Name若为None框架会依据模板模式、prefill/decode 判定、layout、数据类型、B/N/S/D/K 等自动生成用例名见 test_quant_sparse_flash_mla_single.py形如quantSparseFlashMla_CSA_decode_TND_HIF8_PA_BBND_HIF8_1_64_1_1_8192_512_512_000000。执行单用例在 pytest 文件夹路径下执行bash test_run.sh single脚本内部等价于运行python3 -m pytest -rA -s test_quant_sparse_flash_mla_single.py -v -m ci \ -W ignore::UserWarning -W ignore::DeprecationWarning --show-captureno-m ci表示只运行标记为ci的用例。每条用例的执行流水线为笛卡尔积展开参数 →fill_none_params填充默认值 →check_valid_param合法性校验非法则 skip→ golden 数据生成 → eager 直调或 aclgraph 图模式获取 NPU 结果 → 分别校验主输出与 LSE → 记录结果到 Excel默认result.xlsx。单用例模式还支持两个额外选项bash test_run.sh single --run-graph # 切换到 aclgraph 图模式torch.compile torchair bash test_run.sh single --batch-consistency on # 显式开启 batch 一致性校验auto/on/off批量测试批量测试以 Excel 为用例入口支持三种模式均可在命令行指定 Excel 文件、sheet 名与 pt 保存/读取目录命令功能batch_save从 excel 读取用例 → golden 计算 → 保存 pt 文件batch_exec读取 pt 文件 → NPU 测试需先执行batch_savebatch全流程batch_savebatch_exec默认完成后清理 pt 文件选项参数选项说明默认值适用命令--excel 路径指定 excel 文件路径./excel/example.xlsxbatch_save / batch_exec / batch--sheet 名称指定 excel sheet 名decodebatch_save / batch--pt-dir 目录指定 pt 文件保存/读取目录qsmla_testcasebatch_save / batch_exec / batch--keep-pt执行完成后保留 pt 文件默认清理仅 batch--run-graph执行 aclgraph 图模式默认不开启single / batch_exec / batch--batch-consistency auto\|on\|offbatch 一致性策略autosingle / batch_save / batch_exec / batch脚本对参数做了严格校验--keep-pt仅适用于batch命令--run-graph不适用于batch_save未知选项直接报错并打印帮助。这些逻辑定义在 test_run.sh 的参数解析区。使用示例# batch_save只生成 pt 文件 bash test_run.sh batch_save bash test_run.sh batch_save --excel my.xlsx --sheet prefill --pt-dir my_pt # batch_exec只执行 NPU 测试需先 batch_save 生成 pt 文件 bash test_run.sh batch_exec bash test_run.sh batch_exec --pt-dir my_pt # batch全流程默认清理 pt 文件 bash test_run.sh batch # batch全流程保留 pt 文件 bash test_run.sh batch --keep-pt # batch全流程自定义所有参数 bash test_run.sh batch --excel my.xlsx --sheet prefill --pt-dir my_pt --keep-pt从脚本实现看batch_exec在执行前会先检查 pt 目录是否存在且包含.pt文件不存在或为空时报错并退出batch全流程结束后依据KEEP_PT决定是否rm -rf清理 pt 目录。批量执行时通过环境变量传递配置QSMLA_EXCEL、QSMLA_SHEET、QSMLA_PT_DIR、RUN_GRAPH、QSMLA_BATCH_CONSISTENCY、QSMLA_RESULT_SAVE_PATH结果 Excel 路径默认result.xlsx。精度对比机制精度对比是这套框架的核心实现在 result_compare_method.py 的check_result函数中其判定策略如下逐元素比较将 NPU 输出与 CPU golden 展平后逐元素比较支持 bfloat16、fp8e4m3fn/e5m2含 NaN/Inf 位级处理与普通浮点类型阈值设定默认rtol0.005、atol0.000025当结果为 bfloat16 时放宽为rtol0.00781252⁻⁷、atol0.0001通过标准满足np.isclose的元素占比 ≥ 99.5%pct_thd 0.005且最大相对误差小于阈值max_diff_hd 10才判为Pass报告输出打印Rtol/Atol/PctThd/PctRlt/Result摘要失败时输出差异明细前 10 个、后 10 个及最大误差位置双输出校验主输出与 LSEreturn_softmax_lseTrue时分别对比、分开判定任一失败都会反映在最终结果与通过率中取两者最小通过率。单用例与批量测试主程序均复用该函数结果为PASS时用例通过否则以pytest.fail报告失败信息及实际精度。源码级原理golden 复现与 NPU 调用CPU 侧 golden 实现quant_sparse_flash_mla_golden.py 中的GeneralizedSFAQuant类完整复现了算子的前向计算核心逻辑在calculate_by_bnsd稀疏索引收集按template_run_mode分支处理——CSA用gather_cmp_kv按 topk 索引收集压缩 KVHCA用mask_cmp_kv做掩码截取ORI_SPARSE/ORI_CMP_SPARSE用gather_ori_kv收集原始 KV支持-1填充位提前终止分块迭代以s2_base_size 128为分块对 ori_kv 与 cmp_kv 两部分分别做 online softmax 累加量化对齐MM1 结果按softmax_scale * q_descale * kv_descale缩放与 NPU 保持一致注意力分数做 hif8 量化hifp8_scale_value 16.0借助generate_hifloat8_data.py的trans_float_tensor_to_hifuint8/trans_hifuint8_tensor_to_floatMM2 结果再乘kv_descale反量化保证与 NPU 的量化中间过程一致布局转换trans_shape_to_bnsd将 BSND/TND 输入统一转为 BNSD 内部计算格式TND 依赖cu_seqlens_q与seqused_q做 batch 切分计算完成后再由trans_bnsd_to_target_layout转回原布局LSE 输出同样做布局还原。CPU 侧还负责生成稀疏索引数据gen_sparse_indices_bsnd/gen_sparse_indices_tnd依据 mask 模式计算每个 token 的有效 KV 区间causal 阈值、滑动窗口左右边界在有效区间内用randperm生成 topk 索引不足 K 个的部分以-1填充并支持fullK/random两种 topk 长度生成模式。NPU 侧调用链batch/quant_sparse_flash_mla_process.py 提供两种执行入口test_qsmla_quant_process_cieager 直调先构造 metadata 张量调用torch.ops.cann_ops_transformer.quant_sparse_flash_mla_metadata传入头数、头维、cu_seqlens、seqused、topk、cmp_ratio、mask 模式等再直接调用torch.ops.cann_ops_transformer.quant_sparse_flash_mla获取(npu_result, npu_lse)test_qsmla_quant_process_graphaclgraph 图模式将算子封装进torch.nn.Module使用torch.compile torchair 后端CompilerConfig配置reduce-overhead模式、_aclnn_static_shape_kernel静态 shape 内核等编译后执行验证图模式下的行为一致性。两种入口返回(npu_result, cpu_quant_result, cpu_lse, npu_lse)四元组交由测试主程序对比。注意PA_BBND布局在 metadata 中会被映射为PA_ND再传给 metadata 算子。批量一致性校验框架还支持确定性三级 batch 一致性batch_consistency校验通过batch_consistency_seed固定随机种子、batch_consistency_order指定 batch 重排顺序、batch_consistency_batch_split/token_split指定切分方式、batch_consistency_shape_change指定形状变化验证不同执行调度下结果的一致性。该能力复用sparse_flash_mla/tests/pytest/batch_consistency目录下的公共框架见 test_quant_sparse_flash_mla_single.py 的导入通过--batch-consistency auto|on|off控制pytest.ini中consistency标记对应此类用例。小结这套 pytest 测试框架为quant_sparse_flash_mla提供了从参数校验、CPU golden 生成、NPU eager/图模式双通道执行到精度对比与批量一致性校验的完整闭环单用例模式适合算子开发阶段的快速迭代调测批量模式适合回归测试与大规模参数扫描。理解template_run_mode、mask 模式、布局与量化参数之间的约束关系见 check_valid_param.py是高效编写有效用例的关键而 quant_sparse_flash_mla_golden.py 中与 NPU 对齐的量化/反量化细节则是保证精度对比可信度的根基。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考