已合并
fix: 适配sparse_flash_mla算子pytest aclgraph #8759
SH_jingsong创建于 7月15日
fix: 适配sparse_flash_mla算子pytest aclgraph #8759
已合并
SH_jingsong创建于 7月15日
已删除 :master合入到cann/ops-transformermaster
3 个文件变更+466-2
@@ -18,6 +18,8 @@ import pytest
18import random18import random
19import torch19import torch
20import torch_npu20import torch_npu
21+import torchair
22+from torchair.configs.compiler_config import CompilerConfig
21 23 
22# Register sparse_flash_mla and sparse_flash_mla_metadata via PTA24# Register sparse_flash_mla and sparse_flash_mla_metadata via PTA
23TORCH_EXT_PATH = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)),25TORCH_EXT_PATH = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)),
@@ -31,6 +33,43 @@ class Network(torch.nn.Module):
31 def __init__(self):33 def __init__(self):
32 super(Network, self).__init__()34 super(Network, self).__init__()
33 35 
36+ def forward(self, q, ori_kv, cmp_kv, ori_sparse_indices, cmp_sparse_indices,
37+ ori_block_table, cmp_block_table, cu_seqlens_q, cu_seqlens_ori_kv,
38+ cu_seqlens_cmp_kv, seqused_q, seqused_ori_kv, seqused_cmp_kv,
39+ cmp_residual_kv, ori_topk_length, cmp_topk_length, sinks, metadata,
40+ softmax_scale, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left,
41+ ori_win_right, layout_q, layout_kv, topk_value_mode, return_softmax_lse):
42+ npu_result, softmax_lse = torch.ops.cann_ops_transformer.sparse_flash_mla(
43+ q,
44+ ori_kv=ori_kv,
45+ cmp_kv=cmp_kv,
46+ ori_sparse_indices=ori_sparse_indices,
47+ cmp_sparse_indices=cmp_sparse_indices,
48+ ori_block_table=ori_block_table,
49+ cmp_block_table=cmp_block_table,
50+ cu_seqlens_q=cu_seqlens_q,
51+ cu_seqlens_ori_kv=cu_seqlens_ori_kv,
52+ cu_seqlens_cmp_kv=cu_seqlens_cmp_kv,
53+ seqused_q=seqused_q,
54+ seqused_ori_kv=seqused_ori_kv,
55+ seqused_cmp_kv=seqused_cmp_kv,
56+ cmp_residual_kv=cmp_residual_kv,
57+ ori_topk_length=ori_topk_length,
58+ cmp_topk_length=cmp_topk_length,
59+ sinks=sinks,
60+ metadata=metadata,
61+ softmax_scale=softmax_scale,
62+ cmp_ratio=cmp_ratio,
63+ ori_mask_mode=ori_mask_mode,
64+ cmp_mask_mode=cmp_mask_mode,
65+ ori_win_left=ori_win_left,
66+ ori_win_right=ori_win_right,
67+ layout_q=layout_q,
68+ layout_kv=layout_kv,
69+ topk_value_mode=topk_value_mode,
70+ return_softmax_lse=return_softmax_lse)
71+ return npu_result, softmax_lse
72+ 
34def call_npu(input_data):73def call_npu(input_data):
35 params = input_data['params']74 params = input_data['params']
36 metadata_input = input_data['metadata_input']75 metadata_input = input_data['metadata_input']
@@ -154,3 +193,143 @@ def call_npu(input_data):
154 193 
155 torch.npu.synchronize()194 torch.npu.synchronize()
156 return npu_result, softmax_lse195 return npu_result, softmax_lse
196+ 
197+def call_npu_graph(input_data, device_id=0):
198+ params = input_data['params']
199+ metadata_input = input_data['metadata_input']
200+ tensor_input = input_data['input']
201+ print("用例参数(Graph模式): ", params)
202+ torch_npu.npu.set_device(device_id)
203+ 
204+ torch._dynamo.reset()
205+ npu_mode = Network().npu()
206+ config = CompilerConfig()
207+ config.mode = "reduce-overhead"
208+ config.experimental_config.aclgraph._aclnn_static_shape_kernel = True
209+ config.experimental_config.aclgraph._aclnn_static_shape_kernel_build_dir = "./"
210+ config.experimental_config.frozen_parameter = True
211+ config.experimental_config.tiling_schedule_optimize = True
212+ config.experimental_config.topology_sorting_strategy = "StableRDFS"
213+ npu_backend = torchair.get_npu_backend(compiler_config=config)
214+ npu_mode = torch.compile(npu_mode, fullgraph=True, backend=npu_backend, dynamic=False)
215+ 
216+ # metadata解析
217+ K = metadata_input['K']
218+ cmp_ratio = metadata_input['cmp_ratio']
219+ N1 = metadata_input['N1']
220+ N2 = metadata_input['N2']
221+ D = metadata_input['D']
222+ B = metadata_input['B']
223+ 
224+ # tensor解析
225+ q = tensor_input['q'].npu()
226+ ori_kv = tensor_input['ori_kv'].npu() if tensor_input['ori_kv'] is not None else None
227+ cmp_kv = tensor_input['cmp_kv'].npu() if tensor_input['cmp_kv'] is not None else None
228+ ori_block_table = tensor_input['ori_block_table'].npu() if tensor_input['ori_block_table'] is not None else None
229+ cu_seqlens_q = tensor_input['cu_seqlens_q'].npu() if 'cu_seqlens_q' in tensor_input and tensor_input['cu_seqlens_q'] is not None else None
230+ cu_seqlens_ori_kv = tensor_input['cu_seqlens_ori_kv'].npu() if tensor_input['cu_seqlens_ori_kv'] is not None else None
231+ cu_seqlens_cmp_kv = tensor_input['cu_seqlens_cmp_kv'].npu() if tensor_input['cu_seqlens_cmp_kv'] is not None else None
232+ used_seqused_q_flag = False
233+ if 'seqused_q' in tensor_input and tensor_input['seqused_q'] is not None:
234+ seqused_q = tensor_input['seqused_q'].npu()
235+ used_seqused_q_flag = True
236+ else:
237+ seqused_q = None
238+ seqused_ori_kv = tensor_input['seqused_ori_kv'].npu() if tensor_input['seqused_ori_kv'] is not None else None
239+ seqused_cmp_kv = tensor_input['seqused_cmp_kv'].npu() if tensor_input['seqused_cmp_kv'] is not None else None
240+ cmp_residual_kv = tensor_input['cmp_residual_kv'].npu() if tensor_input['cmp_residual_kv'] is not None else None
241+ sinks = tensor_input['sinks'].npu()
242+ softmax_scale = tensor_input['softmax_scale']
243+ ori_mask_mode = tensor_input['ori_mask_mode']
244+ cmp_mask_mode = tensor_input['cmp_mask_mode']
245+ ori_win_left = tensor_input['ori_win_left']
246+ ori_win_right = tensor_input['ori_win_right']
247+ layout_q = tensor_input['layout_q'] if type(tensor_input['layout_q']) == type('TND') else tensor_input['layout_q'][0]
248+ layout_kv = tensor_input['layout_kv']
249+ max_seqlen_q = metadata_input['max_seqlen_q']
250+ max_seqlen_ori_kv = metadata_input['max_seqlen_ori_kv']
251+ max_seqlen_cmp_kv = metadata_input['max_seqlen_cmp_kv']
252+ ori_sparse_indices = tensor_input['ori_sparse_indices']
253+ cmp_sparse_indices = tensor_input['cmp_sparse_indices']
254+ cmp_block_table = tensor_input['cmp_block_table']
255+ ori_topk_length = None
256+ cmp_topk_length = None
257+ return_softmax_lse = params.get('return_softmax_lse')
258+ 
259+ # 将需要上NPU的tensor搬到NPU
260+ if ori_sparse_indices is not None:
261+ ori_sparse_indices = ori_sparse_indices.npu()
262+ if cmp_sparse_indices is not None:
263+ cmp_sparse_indices = cmp_sparse_indices.npu()
264+ if cmp_block_table is not None:
265+ cmp_block_table = cmp_block_table.npu()
266+ 
267+ # 生成 metadata
268+ print("sparse_flash_mla_metadata...")
269+ metadata = torch.ops.cann_ops_transformer.sparse_flash_mla_metadata(
270+ num_heads_q=N1,
271+ num_heads_kv=N2,
272+ head_dim=D,
273+ cu_seqlens_q=cu_seqlens_q,
274+ cu_seqlens_ori_kv=cu_seqlens_ori_kv,
275+ cu_seqlens_cmp_kv=cu_seqlens_cmp_kv,
276+ seqused_q=seqused_q,
277+ seqused_ori_kv=seqused_ori_kv,
278+ seqused_cmp_kv=seqused_cmp_kv,
279+ cmp_residual_kv=cmp_residual_kv,
280+ ori_topk_length=ori_topk_length,
281+ cmp_topk_length=cmp_topk_length,
282+ batch_size=B,
283+ max_seqlen_q=max_seqlen_q,
284+ max_seqlen_ori_kv=max_seqlen_ori_kv,
285+ max_seqlen_cmp_kv=max_seqlen_cmp_kv,
286+ ori_topk=K if ori_sparse_indices is not None else 0,
287+ cmp_topk=K if cmp_sparse_indices is not None else 0,
288+ cmp_ratio=cmp_ratio if cmp_ratio is not None else 1,
289+ ori_mask_mode=ori_mask_mode,
290+ cmp_mask_mode=cmp_mask_mode if cmp_mask_mode is not None else 3,
291+ ori_win_left=ori_win_left,
292+ ori_win_right=ori_win_right,
293+ layout_q=layout_q,
294+ layout_kv=layout_kv,
295+ has_ori_kv=ori_kv != None,
296+ has_cmp_kv=cmp_kv != None)
297+ 
298+ torch.npu.synchronize()
299+ metadata.npu()
300+ 
301+ # 通过编译后的Network调用 sparse_flash_mla
302+ print("sparse_flash_mla (Graph模式)...")
303+ npu_result, softmax_lse = npu_mode(
304+ q=q,
305+ ori_kv=ori_kv,
306+ cmp_kv=cmp_kv,
307+ ori_sparse_indices=ori_sparse_indices,
308+ cmp_sparse_indices=cmp_sparse_indices,
309+ ori_block_table=ori_block_table,
310+ cmp_block_table=cmp_block_table,
311+ cu_seqlens_q=cu_seqlens_q if layout_q == 'TND' else None,
312+ cu_seqlens_ori_kv=cu_seqlens_ori_kv,
313+ cu_seqlens_cmp_kv=cu_seqlens_cmp_kv,
314+ seqused_q=seqused_q if used_seqused_q_flag else None,
315+ seqused_ori_kv=seqused_ori_kv,
316+ seqused_cmp_kv=seqused_cmp_kv,
317+ cmp_residual_kv=cmp_residual_kv,
318+ ori_topk_length=ori_topk_length,
319+ cmp_topk_length=cmp_topk_length,
320+ sinks=sinks,
321+ metadata=metadata,
322+ softmax_scale=softmax_scale,
323+ cmp_ratio=cmp_ratio if cmp_ratio is not None else 1,
324+ ori_mask_mode=ori_mask_mode,
325+ cmp_mask_mode=cmp_mask_mode if cmp_mask_mode is not None else 3,
326+ ori_win_left=ori_win_left,
327+ ori_win_right=ori_win_right,
328+ layout_q=layout_q,
329+ layout_kv=layout_kv,
330+ topk_value_mode=1,
331+ return_softmax_lse=return_softmax_lse if return_softmax_lse is not None else False)
332+ print("sparse_flash_mla (Graph模式) end")
333+ 
334+ torch.npu.synchronize()
335+ return npu_result, softmax_lse
@@ -12,7 +12,7 @@
12# 显示帮助信息12# 显示帮助信息
13show_help() {13show_help() {
14 cat << EOF14 cat << EOF
15-使用方法: $0 {single|save|load} [参数]15+使用方法: $0 {single|save|load|load_graph} [参数]
16 16 
17脚本选项:17脚本选项:
18 single 执行单跑功能18 single 执行单跑功能
@@ -36,6 +36,15 @@ show_help() {
36 示例2: bash $0 load -P "./data" -R "./result/smla_result.xlsx"36 示例2: bash $0 load -P "./data" -R "./result/smla_result.xlsx"
37 示例3: bash $0 load -P "./data" -E "./excel/example.xlsx" -S "CSA"37 示例3: bash $0 load -P "./data" -E "./excel/example.xlsx" -S "CSA"
38 38 
39+ load_graph 执行批量执行PT形式保存的用例的功能(Graph aclgraph模式)
40+ -P pt文件读取地址
41+ -R 结果保存路径
42+ -E excel表地址(指定后仅跑表格中涉及的用例)
43+ -S sheet名(配合-E使用)
44+ 示例1: bash $0 load_graph
45+ 示例2: bash $0 load_graph -P "./data" -R "./result/smla_result.xlsx"
46+ 示例3: bash $0 load_graph -P "./data" -E "./excel/example.xlsx" -S "CSA"
47+ 
39通用选项:48通用选项:
40 -h, --help 显示此帮助信息49 -h, --help 显示此帮助信息
41EOF50EOF
@@ -353,6 +362,175 @@ print('\n'.join(names))
353 echo "结果表格: ${SMLA_RESULT_SAVE_PATH:-./result/smla_result.xlsx}"362 echo "结果表格: ${SMLA_RESULT_SAVE_PATH:-./result/smla_result.xlsx}"
354}363}
355 364 
365+# 运行 批跑执行pt文件(Graph模式) 脚本的函数
366+run_script_load_graph() {
367+ echo "准备运行 test_sparse_flash_mla_batch_graph.py 脚本(Graph模式)..."
368+ 
369+ # 解析 test_sparse_flash_mla_batch_graph.py 的参数
370+ while [[ $# -gt 0 ]]; do
371+ case $1 in
372+ -P)
373+ if [ -z "$2" ] || [[ "$2" == -* ]]; then
374+ echo "错误: -P 参数需要值"
375+ exit 1
376+ fi
377+ P_VALUE="$2"
378+ shift 2
379+ ;;
380+ -R)
381+ if [ -z "$2" ] || [[ "$2" == -* ]]; then
382+ echo "错误: -R 参数需要值"
383+ exit 1
384+ fi
385+ R_VALUE="$2"
386+ shift 2
387+ ;;
388+ -E)
389+ if [ -z "$2" ] || [[ "$2" == -* ]]; then
390+ echo "错误: -E 参数需要值"
391+ exit 1
392+ fi
393+ E_VALUE="$2"
394+ shift 2
395+ ;;
396+ -S)
397+ if [ -z "$2" ] || [[ "$2" == -* ]]; then
398+ echo "错误: -S 参数需要值"
399+ exit 1
400+ fi
401+ S_VALUE="$2"
402+ shift 2
403+ ;;
404+ *)
405+ echo "错误: 未知参数 '$1'"
406+ echo "test_sparse_flash_mla_batch_graph.py 脚本支持的参数: -P, -R, -E, -S"
407+ exit 1
408+ ;;
409+ esac
410+ done
411+ 
412+ # 打印参数信息
413+ echo "==============================="
414+ echo "脚本: test_sparse_flash_mla_batch_graph.py (Graph模式)"
415+ echo "参数配置:"
416+ 
417+ if [ -n "$P_VALUE" ]; then
418+ echo " PT文件读取地址 $P_VALUE"
419+ export SMLA_PT_LOAD_PATH="$P_VALUE"
420+ else
421+ echo " 默认PT文件读取地址 ./data"
422+ fi
423+ 
424+ if [ -n "$R_VALUE" ]; then
425+ echo " 结果存储至文件 $R_VALUE"
426+ export SMLA_RESULT_SAVE_PATH="$R_VALUE"
427+ else
428+ echo " 结果存储至文件 ./result/smla_result.xlsx"
429+ fi
430+ 
431+ if [ -n "$E_VALUE" ]; then
432+ echo " Excel文件路径 $E_VALUE"
433+ export SMLA_EXCEL_PATH="$E_VALUE"
434+ export SMLA_BATCH_TEST_MODE=1
435+ if [ -n "$S_VALUE" ]; then
436+ echo " 使用sheet名 $S_VALUE"
437+ export SMLA_EXCEL_SHEET="$S_VALUE"
438+ fi
439+ else
440+ echo " 未指定Excel,全量批跑目录下所有.pt文件"
441+ fi
442+ 
443+ echo "==============================="
444+ 
445+ # 检查脚本是否存在
446+ if [ ! -f "test_sparse_flash_mla_batch_graph.py" ]; then
447+ echo "错误: 找不到 test_sparse_flash_mla_batch_graph.py 脚本"
448+ exit 1
449+ fi
450+ 
451+ # 获取用例目录
452+ LOAD_DIR="${SMLA_PT_LOAD_PATH:-./data}"
453+ if [ ! -d "$LOAD_DIR" ]; then
454+ echo "错误: 用例目录不存在: $LOAD_DIR"
455+ exit 1
456+ fi
457+ 
458+ if [ "${SMLA_BATCH_TEST_MODE}" = "1" ] && [ -n "$E_VALUE" ]; then
459+ SHEET_NAME="${S_VALUE:-CSA}"
460+ echo "按表格筛选模式: Excel=$E_VALUE, Sheet=$SHEET_NAME"
461+ TARGET_NAMES=$(python3 -c "
462+import pandas as pd
463+df = pd.read_excel('$E_VALUE', sheet_name='$SHEET_NAME')
464+names = [str(n) for n in df['testcase_name'].dropna().tolist() if str(n) != 'None']
465+print('\n'.join(names))
466+")
467+ if [ -z "$TARGET_NAMES" ]; then
468+ echo "错误: 表格中没有有效的testcase_name"
469+ exit 1
470+ fi
471+ CASE_FILES=()
472+ while IFS= read -r target_name; do
473+ while IFS= read -r matched_file; do
474+ if [ -n "$matched_file" ] && [[ ! " ${CASE_FILES[*]} " =~ " ${matched_file} " ]]; then
475+ CASE_FILES+=("$matched_file")
476+ fi
477+ done < <(find "$LOAD_DIR" -maxdepth 1 -name "*.pt" | grep "$target_name" | sort)
478+ done <<< "$TARGET_NAMES"
479+ echo "从表格中读取到目标用例名, 筛选后共 ${#CASE_FILES[@]} 个.pt文件"
480+ else
481+ mapfile -t CASE_FILES < <(find "$LOAD_DIR" -maxdepth 1 -name "*.pt" | sort)
482+ fi
483+ 
484+ TOTAL=${#CASE_FILES[@]}
485+ if [ "$TOTAL" -eq 0 ]; then
486+ echo "错误: 目录 $LOAD_DIR 下未找到匹配的 .pt 用例"
487+ exit 1
488+ fi
489+ 
490+ echo "共 $TOTAL 条用例待执行, 目录: $LOAD_DIR"
491+ echo "开始隔离批量执行(Graph模式, 每条用例独立 pytest 进程)..."
492+ 
493+ PASS=0
494+ FAIL=0
495+ FAIL_LIST=()
496+ SUMMARY_LOG="batch_graph_summary.log"
497+ FAIL_LOG="batch_graph_fail_list.log"
498+ : > "$SUMMARY_LOG"
499+ : > "$FAIL_LOG"
500+ 
501+ i=0
502+ for case_file in "${CASE_FILES[@]}"; do
503+ i=$((i+1))
504+ case_name=$(basename "$case_file")
505+ echo -e "\n===== [$i/$TOTAL] 执行用例(Graph模式): $case_name =====" | tee -a "$SUMMARY_LOG"
506+ 
507+ QSAS_TESTCASE_PATH="$case_file" python3 -m pytest -rA -s test_sparse_flash_mla_batch_graph.py -v -m graph 2>&1 | tee -a "$SUMMARY_LOG"
508+ status=${PIPESTATUS[0]}
509+ 
510+ if [ "$status" -eq 0 ]; then
511+ PASS=$((PASS+1))
512+ echo "[PASS] $case_name" | tee -a "$SUMMARY_LOG"
513+ else
514+ FAIL=$((FAIL+1))
515+ FAIL_LIST+=("$case_name")
516+ echo "[FAIL] $case_name" | tee -a "$SUMMARY_LOG"
517+ echo "$case_name" >> "$FAIL_LOG"
518+ fi
519+ done
520+ 
521+ echo -e "\n========== 批量执行汇总(Graph模式) ==========" | tee -a "$SUMMARY_LOG"
522+ echo "总计: $TOTAL 通过: $PASS 失败: $FAIL" | tee -a "$SUMMARY_LOG"
523+ if [ "$FAIL" -gt 0 ]; then
524+ echo "失败用例:" | tee -a "$SUMMARY_LOG"
525+ for f in "${FAIL_LIST[@]}"; do
526+ echo " - $f" | tee -a "$SUMMARY_LOG"
527+ done
528+ fi
529+ echo "详细日志: $SUMMARY_LOG"
530+ echo "失败清单: $FAIL_LOG"
531+ echo "结果表格: ${SMLA_RESULT_SAVE_PATH:-./result/smla_result.xlsx}"
532+}
533+ 
356# 根据脚本类型调用相应的函数534# 根据脚本类型调用相应的函数
357case "$SCRIPT_TYPE" in535case "$SCRIPT_TYPE" in
358 single)536 single)
@@ -364,9 +542,12 @@ case "$SCRIPT_TYPE" in
364 load)542 load)
365 run_script_load "$@"543 run_script_load "$@"
366 ;;544 ;;
545+ load_graph)
546+ run_script_load_graph "$@"
547+ ;;
367 *)548 *)
368 echo "错误: 未知的脚本类型 '$SCRIPT_TYPE'"549 echo "错误: 未知的脚本类型 '$SCRIPT_TYPE'"
369- echo "可用类型: single, save, load"550+ echo "可用类型: single, save, load, load_graph"
370 show_help551 show_help
371 exit 1552 exit 1
372 ;;553 ;;
@@ -0,0 +1,104 @@
1+#!/usr/bin/python
2+# -*- coding: utf-8 -*-
3+# -----------------------------------------------------------------------------------------------------------
4+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
5+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
6+# CANN Open Software License Agreement Version 2.0 (the "License").
7+# Please refer to the License for details. You may not use this file except in compliance with the License.
8+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
9+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
10+# See LICENSE in the root of the software repository for the full text of the License.
11+# -----------------------------------------------------------------------------------------------------------
12+ 
13+import torch
14+import torch_npu
15+import result_compare_method
16+import utils
17+from batch import sparse_flash_mla_process
18+import pytest
19+import concurrent.futures
20+import pandas as pd
21+from pathlib import Path
22+import os
23+ 
24+pt_dir = os.getenv("SMLA_PT_LOAD_PATH", "./data")
25+result_path = Path(os.getenv("SMLA_RESULT_SAVE_PATH", './result/smla_result.xlsx'))
26+batch_test_mode = int(os.environ.get("SMLA_BATCH_TEST_MODE", 0))
27+excel_path = os.environ.get("SMLA_EXCEL_PATH", os.path.join(os.path.dirname(__file__), "excel", "example.xlsx"))
28+excel_sheet = os.environ.get("SMLA_EXCEL_SHEET", "CSA")
29+device_id = int(os.environ.get("SMLA_DEVICE_ID", 0))
30+ 
31+_single_case_path = os.environ.get("QSAS_TESTCASE_PATH", "").strip()
32+ 
33+locals()["testcase_files"] = []
34+if _single_case_path:
35+ if not os.path.isfile(_single_case_path):
36+ print(f"错误: 环境变量 QSAS_TESTCASE_PATH 指定的用例文件不存在: {_single_case_path}")
37+ else:
38+ print(f"单用例隔离模式, 仅执行: {_single_case_path}")
39+ locals()["testcase_files"].append(_single_case_path)
40+elif os.path.isdir(pt_dir):
41+ pt_files = [f for f in os.listdir(pt_dir) if f.endswith('.pt')]
42+ if not pt_files:
43+ print(f"错误: 目录中没有找到.pt文件: {pt_dir}")
44+ elif batch_test_mode == 1:
45+ df = pd.read_excel(excel_path, sheet_name=excel_sheet)
46+ target_names = [str(name) for name in df['testcase_name'].dropna().tolist() if str(name) != 'None']
47+ if not target_names:
48+ print(f"错误: 表格中没有有效的testcase_name: {excel_path} sheet: {excel_sheet}")
49+ else:
50+ print(f"从表格[{excel_sheet}]中读取到 {len(target_names)} 个目标用例名")
51+ for target_name in target_names:
52+ matched = [f for f in pt_files if target_name in f]
53+ if matched:
54+ for f in matched:
55+ filepath = os.path.join(pt_dir, f)
56+ if filepath not in locals()["testcase_files"]:
57+ locals()["testcase_files"].append(filepath)
58+ else:
59+ print(f"警告: 用例名 '{target_name}' 未匹配到任何.pt文件")
60+ print(f"按表格筛选后共 {len(locals()['testcase_files'])} 个测试用例文件")
61+ else:
62+ print(f"找到 {len(pt_files)} 个测试用例文件")
63+ for pt_file in pt_files:
64+ filepath = os.path.join(pt_dir, pt_file)
65+ locals()["testcase_files"].append(filepath)
66+else:
67+ print(f"错误: 输出目录不存在: {pt_dir}")
68+ 
69+print("files:", locals()["testcase_files"])
70+ 
71+def smla_graph(testcase_files):
72+ test_data = torch.load(testcase_files, map_location="cpu")
73+ npu_error_msg = None
74+ try:
75+ npu_result, softmax_lse = sparse_flash_mla_process.call_npu_graph(test_data, device_id=device_id)
76+ result, fulfill_percent = result_compare_method.check_result(test_data['cpu_output'], npu_result)
77+ if test_data['params'].get('return_softmax_lse'):
78+ print("return_softmax_lse is true!!!")
79+ result, fulfill_percent = result_compare_method.check_result(test_data['softmax_lse'], softmax_lse)
80+ except Exception as e:
81+ npu_error_msg = str(e)
82+ print("NPU ERROR:", npu_error_msg)
83+ result = "NPU ERROR"
84+ fulfill_percent = 0
85+ 
86+ utils.save_result(result, fulfill_percent, test_data['params'], result_path)
87+ 
88+ if result == "Failed":
89+ pytest.fail(f"用例精度失败:{os.path.basename(testcase_files)} 精度:{fulfill_percent:.2f}%")
90+ if result == "NPU ERROR":
91+ pytest.fail(f"用例执行失败:{os.path.basename(testcase_files)} NPU ERROR: {npu_error_msg}")
92+ 
93+testcase_ids = [os.path.splitext(os.path.basename(f))[0] for f in locals()["testcase_files"]]
94+ 
95+@pytest.mark.graph
96+@pytest.mark.parametrize("testcase_files", locals()["testcase_files"], ids=testcase_ids)
97+def test_sparse_flash_mla(testcase_files):
98+ with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:
99+ futures = executor.submit(smla_graph, testcase_files)
100+ for future in concurrent.futures.as_completed([futures]):
101+ try:
102+ result = future.result()
103+ except Exception as e:
104+ pytest.fail(f"当前用例线程执行失败")