已合并
fix: 适配sparse_flash_mla算子pytest aclgraph #8759
SH_jingsong创建于 7月15日
fix: 适配sparse_flash_mla算子pytest aclgraph #8759
已合并
从已删除 :master合入到cann/ops-transformermaster
共 3 个文件变更+466-2
| @@ -18,6 +18,8 @@ import pytest | |||
| 18 | import random | 18 | import random |
| 19 | import torch | 19 | import torch |
| 20 | import torch_npu | 20 | import 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 PTA | 24 | # Register sparse_flash_mla and sparse_flash_mla_metadata via PTA |
| 23 | TORCH_EXT_PATH = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)), | 25 | TORCH_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 | + | ||
| 34 | def call_npu(input_data): | 73 | def 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_lse | 195 | 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 | # 显示帮助信息 |
| 13 | show_help() { | 13 | show_help() { |
| 14 | cat << EOF | 14 | 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 显示此帮助信息 |
| 41 | EOF | 50 | EOF |
| @@ -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 | # 根据脚本类型调用相应的函数 |
| 357 | case "$SCRIPT_TYPE" in | 535 | case "$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_help | 551 | show_help |
| 371 | exit 1 | 552 | 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 | + | ||
| 96 | + | ||
| 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"当前用例线程执行失败") | ||