已合并
feat: add matrix_set_diag_v2 UT #4608
zhanw_coding创建于 8月13日
feat: add matrix_set_diag_v2 UT #4608
已合并
共 5 个文件变更+423-2
| @@ -938,6 +938,9 @@ VC5@ops-math: | |||
| 938 | # depth_to_space | 938 | # depth_to_space |
| 939 | - ops/ops-math/conversion/depth_to_space/op_host/ | 939 | - ops/ops-math/conversion/depth_to_space/op_host/ |
| 940 | - ops/ops-math/conversion/depth_to_space/op_kernel/ | 940 | - ops/ops-math/conversion/depth_to_space/op_kernel/ |
| 941 | + # matrix_set_diag_v2 | ||
| 942 | + - ops/ops-math/conversion/matrix_set_diag_v2/op_host/ | ||
| 943 | + - ops/ops-math/conversion/matrix_set_diag_v2/op_kernel/ | ||
| 941 | # mirror_pad | 944 | # mirror_pad |
| 942 | - ops/ops-math/conversion/mirror_pad/op_host/ | 945 | - ops/ops-math/conversion/mirror_pad/op_host/ |
| 943 | - ops/ops-math/conversion/mirror_pad/op_kernel/ | 946 | - ops/ops-math/conversion/mirror_pad/op_kernel/ |
| @@ -1010,6 +1013,9 @@ VC5@ops-math: | |||
| 1010 | - ops/ops-math/conversion/depth_to_space/examples/*.cpp | 1013 | - ops/ops-math/conversion/depth_to_space/examples/*.cpp |
| 1011 | - ops/ops-math/conversion/depth_to_space/tests/**/*.cpp | 1014 | - ops/ops-math/conversion/depth_to_space/tests/**/*.cpp |
| 1012 | - ops/ops-math/conversion/depth_to_space/tests/**/*.py | 1015 | - ops/ops-math/conversion/depth_to_space/tests/**/*.py |
| 1016 | + # matrix_set_diag_v2 | ||
| 1017 | + - ops/ops-math/conversion/matrix_set_diag_v2/examples/ | ||
| 1018 | + - ops/ops-math/conversion/matrix_set_diag_v2/tests/ | ||
| 1013 | # mirror_pad | 1019 | # mirror_pad |
| 1014 | - ops/ops-math/conversion/mirror_pad/examples/ | 1020 | - ops/ops-math/conversion/mirror_pad/examples/ |
| 1015 | - ops/ops-math/conversion/mirror_pad/tests/ | 1021 | - ops/ops-math/conversion/mirror_pad/tests/ |
| @@ -1096,6 +1102,9 @@ VC5@ops-math: | |||
| 1096 | - ops/ops-math/conversion/pad_v2/op_kernel/ | 1102 | - ops/ops-math/conversion/pad_v2/op_kernel/ |
| 1097 | - ops/ops-math/conversion/pad_v2/op_host/arch35/ | 1103 | - ops/ops-math/conversion/pad_v2/op_host/arch35/ |
| 1098 | - ops/ops-math/conversion/pad_v2/op_host/pad_v2_def.cpp | 1104 | - ops/ops-math/conversion/pad_v2/op_host/pad_v2_def.cpp |
| 1105 | + - ops/ops-math/conversion/matrix_set_diag_v2/op_kernel/ | ||
| 1106 | + - ops/ops-math/conversion/matrix_set_diag_v2/op_host/arch35/ | ||
| 1107 | + - ops/ops-math/conversion/matrix_set_diag_v2/op_host/matrix_set_diag_v2_def.cpp | ||
| 1099 | - ops/ops-math/conversion/mirror_pad/op_kernel/ | 1108 | - ops/ops-math/conversion/mirror_pad/op_kernel/ |
| 1100 | - ops/ops-math/conversion/mirror_pad/op_host/arch35/ | 1109 | - ops/ops-math/conversion/mirror_pad/op_host/arch35/ |
| 1101 | - ops/ops-math/conversion/mirror_pad/op_host/mirror_pad_def.cpp | 1110 | - ops/ops-math/conversion/mirror_pad/op_host/mirror_pad_def.cpp |
| @@ -112,8 +112,10 @@ ge::graphStatus MatrixSetDiagTiling::ParamCheck() | |||
| 112 | 112 | ||
| 113 | inputInfo_.xColNum = inputShapeVal.GetDim(dimNum_ - COL_DIM_OFFSET); | 113 | inputInfo_.xColNum = inputShapeVal.GetDim(dimNum_ - COL_DIM_OFFSET); |
| 114 | inputInfo_.xRowNum = inputShapeVal.GetDim(dimNum_ - ROW_DIM_OFFSET); | 114 | inputInfo_.xRowNum = inputShapeVal.GetDim(dimNum_ - ROW_DIM_OFFSET); |
| 115 | - inputInfo_.maxDiagLen = static_cast<size_t>(diagShapeVal.GetDim(diagDimNum_ - 1)); | 115 | + uint64_t maxDiagLen = diagShapeVal.GetDim(diagDimNum_ - 1); |
| 116 | - OP_CHECK_IF(inputInfo_.maxDiagLen != std::min(inputInfo_.xColNum, inputInfo_.xRowNum), | 116 | + inputInfo_.maxDiagLen = static_cast<uint32_t>(maxDiagLen); |
| 117 | + OP_CHECK_IF((inputInfo_.maxDiagLen != std::min(inputInfo_.xColNum, inputInfo_.xRowNum)) || | ||
| 118 | + (inputInfo_.maxDiagLen != maxDiagLen), | ||
| 117 | OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_->GetNodeName(), "diagonal", | 119 | OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_->GetNodeName(), "diagonal", |
| 118 | Ops::Base::ToString(diagShapeVal).c_str(), | 120 | Ops::Base::ToString(diagShapeVal).c_str(), |
| 119 | "diagonal length must equal min(row, col) of input"), | 121 | "diagonal length must equal min(row, col) of input"), |
| @@ -152,6 +152,25 @@ TEST_F(MatrixSetDiagTilingTest, test_tiling_failed_diag_len_invalid) | |||
| 152 | ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED); | 152 | ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED); |
| 153 | } | 153 | } |
| 154 | 154 | ||
| 155 | +TEST_F(MatrixSetDiagTilingTest, test_tiling_failed_diag_len_overflow) | ||
| 156 | +{ | ||
| 157 | + // UINT32_MAX + 4 = 4294967299, static_cast<uint32_t> truncates to 3 | ||
| 158 | + // condition1 (maxDiagLen != min(row, col)): 3 != min(3, 4294967299) = 3 -> FALSE | ||
| 159 | + // condition2 (maxDiagLen_u32 != maxDiagLen_u64): 3 != 4294967299 -> TRUE (truncation detected) | ||
| 160 | + MatrixSetDiagCompileInfo compileInfo = {}; | ||
| 161 | + constexpr int64_t overflowDim = 4294967299LL; | ||
| 162 | + gert::TilingContextPara tilingContextPara("MatrixSetDiag", | ||
| 163 | + { | ||
| 164 | + {{{3, overflowDim}, {3, overflowDim}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 165 | + {{{overflowDim}, {overflowDim}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 166 | + }, | ||
| 167 | + { | ||
| 168 | + {{{3, overflowDim}, {3, overflowDim}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 169 | + }, | ||
| 170 | + &compileInfo); | ||
| 171 | + ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED); | ||
| 172 | +} | ||
| 173 | + | ||
| 155 | TEST_F(MatrixSetDiagTilingTest, test_tiling_failed_batch_dim_invalid) | 174 | TEST_F(MatrixSetDiagTilingTest, test_tiling_failed_batch_dim_invalid) |
| 156 | { | 175 | { |
| 157 | MatrixSetDiagCompileInfo compileInfo = {}; | 176 | MatrixSetDiagCompileInfo compileInfo = {}; |
| @@ -372,3 +372,262 @@ TEST_F(MatrixSetDiagV2TilingTest, test_tiling_cut_tail_optimize_core_num2) | |||
| 372 | EXPECT_EQ(data->xColFactor, 2); | 372 | EXPECT_EQ(data->xColFactor, 2); |
| 373 | EXPECT_EQ(data->totalCntPerCore, 1); | 373 | EXPECT_EQ(data->totalCntPerCore, 1); |
| 374 | } | 374 | } |
| 375 | + | ||
| 376 | +// 场景:k0==k1==0 且 dSize<=2、tailAxisDataSize>=32767,触发 Tiling4CutTail 后走 V1 路径 | ||
| 377 | +// (Tiling4CutW/CalUbFactor/GetOptimizeTiling 循环并在 sizeTaken<=4096 时 break)。 | ||
| 378 | +// 输入 x{10000,4} INT16,diag{4},k{0}。 | ||
| 379 | +// 期望:成功;tilingKey=0x404(V1+切尾轴);tilingData 为 MatrixSetDiagTilingData: | ||
| 380 | +// coreNum=23, mergeDimSize=1, xRowNum=10000, xColNum=4, diagLen=4, ubPerCore=1, | ||
| 381 | +// ubFactor=1740, ubTotalCount=23, ubPerTail=23, tailAxisDataSize=40000;blockNum=23;workspace=0。 | ||
| 382 | +TEST_F(MatrixSetDiagV2TilingTest, test_tiling_v1_cutw_optimize_loop) | ||
| 383 | +{ | ||
| 384 | + std::vector<int32_t> kValues = {0}; | ||
| 385 | + gert::TilingContextPara tilingContextPara("MatrixSetDiagV2", | ||
| 386 | + { | ||
| 387 | + {{{10000, 4}, {10000, 4}}, ge::DT_INT16, ge::FORMAT_ND}, | ||
| 388 | + {{{4}, {4}}, ge::DT_INT16, ge::FORMAT_ND}, | ||
| 389 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 390 | + }, | ||
| 391 | + { | ||
| 392 | + {{{10000, 4}, {10000, 4}}, ge::DT_INT16, ge::FORMAT_ND}, | ||
| 393 | + }, | ||
| 394 | + &compileInfo); | ||
| 395 | + TilingInfo tilingInfo; | ||
| 396 | + auto tilingRet = ExecuteTiling(tilingContextPara, tilingInfo); | ||
| 397 | + ASSERT_TRUE(tilingRet); | ||
| 398 | + std::vector<int64_t> expectWorkspaces = {0}; | ||
| 399 | + EXPECT_EQ(tilingInfo.workspaceSizes, expectWorkspaces); | ||
| 400 | + EXPECT_EQ(tilingInfo.blockNum, 23); | ||
| 401 | + EXPECT_EQ(tilingInfo.tilingKey, 0b1'0'0'00000100); | ||
| 402 | + ASSERT_EQ(tilingInfo.tilingDataSize, sizeof(MatrixSetDiagTilingData)); | ||
| 403 | + MatrixSetDiagTilingData* data = reinterpret_cast<MatrixSetDiagTilingData*>(tilingInfo.tilingData.get()); | ||
| 404 | + EXPECT_EQ(data->coreNum, 23); | ||
| 405 | + EXPECT_EQ(data->mergeDimSize, 1); | ||
| 406 | + EXPECT_EQ(data->xRowNum, 10000); | ||
| 407 | + EXPECT_EQ(data->xColNum, 4); | ||
| 408 | + EXPECT_EQ(data->diagLen, 4); | ||
| 409 | + EXPECT_EQ(data->ubPerCore, 1); | ||
| 410 | + EXPECT_EQ(data->ubFactor, 1740); | ||
| 411 | + EXPECT_EQ(data->ubTotalCount, 23); | ||
| 412 | + EXPECT_EQ(data->ubPerTail, 23); | ||
| 413 | + EXPECT_EQ(data->tailAxisDataSize, 40000); | ||
| 414 | +} | ||
| 415 | + | ||
| 416 | +// 场景:k0==k1==0 且 dSize<=2、tailAxisDataSize>=32767,V1 路径中 | ||
| 417 | +// realCoreNum(52)/coreNum(64)=0.8125>=MIN_USED_CORES_RATIO(0.8), | ||
| 418 | +// 触发 GetOptimizeTiling 的提前返回分支。输入 x{52,10000,4} INT16,diag{52,4},k{0}。 | ||
| 419 | +// 期望:成功;tilingKey=0x404;tilingData:coreNum=52, mergeDimSize=52, xRowNum=10000, xColNum=4, | ||
| 420 | +// diagLen=4, ubPerCore=1, ubFactor=40000, ubTotalCount=52, ubPerTail=1, tailAxisDataSize=40000; | ||
| 421 | +// blockNum=52;workspace=0。 | ||
| 422 | +TEST_F(MatrixSetDiagV2TilingTest, test_tiling_v1_cutw_optimize_early_return) | ||
| 423 | +{ | ||
| 424 | + std::vector<int32_t> kValues = {0}; | ||
| 425 | + gert::TilingContextPara tilingContextPara("MatrixSetDiagV2", | ||
| 426 | + { | ||
| 427 | + {{{52, 10000, 4}, {52, 10000, 4}}, ge::DT_INT16, ge::FORMAT_ND}, | ||
| 428 | + {{{52, 4}, {52, 4}}, ge::DT_INT16, ge::FORMAT_ND}, | ||
| 429 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 430 | + }, | ||
| 431 | + { | ||
| 432 | + {{{52, 10000, 4}, {52, 10000, 4}}, ge::DT_INT16, ge::FORMAT_ND}, | ||
| 433 | + }, | ||
| 434 | + &compileInfo); | ||
| 435 | + TilingInfo tilingInfo; | ||
| 436 | + auto tilingRet = ExecuteTiling(tilingContextPara, tilingInfo); | ||
| 437 | + ASSERT_TRUE(tilingRet); | ||
| 438 | + std::vector<int64_t> expectWorkspaces = {0}; | ||
| 439 | + EXPECT_EQ(tilingInfo.workspaceSizes, expectWorkspaces); | ||
| 440 | + EXPECT_EQ(tilingInfo.blockNum, 52); | ||
| 441 | + EXPECT_EQ(tilingInfo.tilingKey, 0b1'0'0'00000100); | ||
| 442 | + ASSERT_EQ(tilingInfo.tilingDataSize, sizeof(MatrixSetDiagTilingData)); | ||
| 443 | + MatrixSetDiagTilingData* data = reinterpret_cast<MatrixSetDiagTilingData*>(tilingInfo.tilingData.get()); | ||
| 444 | + EXPECT_EQ(data->coreNum, 52); | ||
| 445 | + EXPECT_EQ(data->mergeDimSize, 52); | ||
| 446 | + EXPECT_EQ(data->xRowNum, 10000); | ||
| 447 | + EXPECT_EQ(data->xColNum, 4); | ||
| 448 | + EXPECT_EQ(data->diagLen, 4); | ||
| 449 | + EXPECT_EQ(data->ubPerCore, 1); | ||
| 450 | + EXPECT_EQ(data->ubFactor, 40000); | ||
| 451 | + EXPECT_EQ(data->ubTotalCount, 52); | ||
| 452 | + EXPECT_EQ(data->ubPerTail, 1); | ||
| 453 | + EXPECT_EQ(data->tailAxisDataSize, 40000); | ||
| 454 | +} | ||
| 455 | + | ||
| 456 | +// 场景:k0==k1==0,xColNum*dSize >= bufferSize_,触发 V1 路径 CalUbFactor 的 if 分支 | ||
| 457 | +// (ubFactor=validBufSize/dSize)。输入 x{2,16384} INT64,diag{2},k{0}。 | ||
| 458 | +// 期望:成功;tilingKey=0x404;tilingData:coreNum=51, mergeDimSize=1, xRowNum=2, xColNum=16384, | ||
| 459 | +// diagLen=2, ubPerCore=1, ubFactor=643, ubTotalCount=51, ubPerTail=51, tailAxisDataSize=32768; | ||
| 460 | +// blockNum=51;workspace=0。 | ||
| 461 | +TEST_F(MatrixSetDiagV2TilingTest, test_tiling_v1_cutw_calubfactor_if_branch) | ||
| 462 | +{ | ||
| 463 | + std::vector<int32_t> kValues = {0}; | ||
| 464 | + gert::TilingContextPara tilingContextPara("MatrixSetDiagV2", | ||
| 465 | + { | ||
| 466 | + {{{2, 16384}, {2, 16384}}, ge::DT_INT64, ge::FORMAT_ND}, | ||
| 467 | + {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND}, | ||
| 468 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 469 | + }, | ||
| 470 | + { | ||
| 471 | + {{{2, 16384}, {2, 16384}}, ge::DT_INT64, ge::FORMAT_ND}, | ||
| 472 | + }, | ||
| 473 | + &compileInfo); | ||
| 474 | + TilingInfo tilingInfo; | ||
| 475 | + auto tilingRet = ExecuteTiling(tilingContextPara, tilingInfo); | ||
| 476 | + ASSERT_TRUE(tilingRet); | ||
| 477 | + std::vector<int64_t> expectWorkspaces = {0}; | ||
| 478 | + EXPECT_EQ(tilingInfo.workspaceSizes, expectWorkspaces); | ||
| 479 | + EXPECT_EQ(tilingInfo.blockNum, 51); | ||
| 480 | + EXPECT_EQ(tilingInfo.tilingKey, 0b1'0'0'00000100); | ||
| 481 | + ASSERT_EQ(tilingInfo.tilingDataSize, sizeof(MatrixSetDiagTilingData)); | ||
| 482 | + MatrixSetDiagTilingData* data = reinterpret_cast<MatrixSetDiagTilingData*>(tilingInfo.tilingData.get()); | ||
| 483 | + EXPECT_EQ(data->coreNum, 51); | ||
| 484 | + EXPECT_EQ(data->mergeDimSize, 1); | ||
| 485 | + EXPECT_EQ(data->xRowNum, 2); | ||
| 486 | + EXPECT_EQ(data->xColNum, 16384); | ||
| 487 | + EXPECT_EQ(data->diagLen, 2); | ||
| 488 | + EXPECT_EQ(data->ubPerCore, 1); | ||
| 489 | + EXPECT_EQ(data->ubFactor, 643); | ||
| 490 | + EXPECT_EQ(data->ubTotalCount, 51); | ||
| 491 | + EXPECT_EQ(data->ubPerTail, 51); | ||
| 492 | + EXPECT_EQ(data->tailAxisDataSize, 32768); | ||
| 493 | +} | ||
| 494 | + | ||
| 495 | +// 场景:ratio<SIMT_RATIO 使 way=SIMT,且 ubSize(16384)<SIMT_DCACHE_SIZE(32768), | ||
| 496 | +// CalculateValidBufSize 的 SIMT 分支校验失败。输入 x{20,20} FLOAT,diag{20},k{0},ubSize=16384。 | ||
| 497 | +// 期望:GRAPH_FAILED。 | ||
| 498 | +TEST_F(MatrixSetDiagV2TilingTest, test_tiling_simt_ub_too_small_fail) | ||
| 499 | +{ | ||
| 500 | + std::vector<int32_t> kValues = {0}; | ||
| 501 | + gert::TilingContextPara tilingContextPara("MatrixSetDiagV2", | ||
| 502 | + { | ||
| 503 | + {{{20, 20}, {20, 20}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 504 | + {{{20}, {20}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 505 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 506 | + }, | ||
| 507 | + { | ||
| 508 | + {{{20, 20}, {20, 20}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 509 | + }, | ||
| 510 | + &compileInfo, 64, 16384); | ||
| 511 | + ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED); | ||
| 512 | +} | ||
| 513 | + | ||
| 514 | +// 场景:ratio=1.99>=SCATTER_RATIO 使 way=GATHER 且 isVLFullLoad=false(additionTileSize=整块 tail 数据), | ||
| 515 | +// validBufSize < totalTailSize 使 ubFactor=0,Tiling4NoCutTail 回退 Tiling4CutTail。 | ||
| 516 | +// 输入 x{100,100} FLOAT,diag{199,100},k{-99,99}。 | ||
| 517 | +// 期望:成功;tilingKey=0x400(回退后切尾轴);tilingData 为 MSDV2CutTailTilingData: | ||
| 518 | +// input.coreNum=3, mergeDimSize=1, xRowNum=100, xColNum=100, diagNum=199, maxDiagLen=100, | ||
| 519 | +// k0=-99, k1=99, xRowFactor=40, xColFactor=100, totalCntPerCore=1;blockNum=3;workspace=0。 | ||
| 520 | +TEST_F(MatrixSetDiagV2TilingTest, test_tiling_gather_ubfactor_zero_fallback_cuttail) | ||
| 521 | +{ | ||
| 522 | + std::vector<int32_t> kValues = {-99, 99}; | ||
| 523 | + gert::TilingContextPara tilingContextPara("MatrixSetDiagV2", | ||
| 524 | + { | ||
| 525 | + {{{100, 100}, {100, 100}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 526 | + {{{199, 100}, {199, 100}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 527 | + {{{2}, {2}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 528 | + }, | ||
| 529 | + { | ||
| 530 | + {{{100, 100}, {100, 100}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 531 | + }, | ||
| 532 | + &compileInfo); | ||
| 533 | + TilingInfo tilingInfo; | ||
| 534 | + auto tilingRet = ExecuteTiling(tilingContextPara, tilingInfo); | ||
| 535 | + ASSERT_TRUE(tilingRet); | ||
| 536 | + std::vector<int64_t> expectWorkspaces = {0}; | ||
| 537 | + EXPECT_EQ(tilingInfo.workspaceSizes, expectWorkspaces); | ||
| 538 | + EXPECT_EQ(tilingInfo.blockNum, 3); | ||
| 539 | + EXPECT_EQ(tilingInfo.tilingKey, 0b1'0'0'00000000); | ||
| 540 | + ASSERT_EQ(tilingInfo.tilingDataSize, sizeof(MSDV2CutTailTilingData)); | ||
| 541 | + MSDV2CutTailTilingData* data = reinterpret_cast<MSDV2CutTailTilingData*>(tilingInfo.tilingData.get()); | ||
| 542 | + EXPECT_EQ(data->input.coreNum, 3); | ||
| 543 | + EXPECT_EQ(data->input.mergeDimSize, 1); | ||
| 544 | + EXPECT_EQ(data->input.xRowNum, 100); | ||
| 545 | + EXPECT_EQ(data->input.xColNum, 100); | ||
| 546 | + EXPECT_EQ(data->input.diagNum, 199); | ||
| 547 | + EXPECT_EQ(data->input.maxDiagLen, 100); | ||
| 548 | + EXPECT_EQ(data->input.k0, -99); | ||
| 549 | + EXPECT_EQ(data->input.k1, 99); | ||
| 550 | + EXPECT_EQ(data->xRowFactor, 40); | ||
| 551 | + EXPECT_EQ(data->xColFactor, 100); | ||
| 552 | + EXPECT_EQ(data->totalCntPerCore, 1); | ||
| 553 | +} | ||
| 554 | + | ||
| 555 | +// 场景:ratio=0.25 使 way=SCATTER 且 isVLFullLoad=true;NoCutTail 优化路径 | ||
| 556 | +// GetOptimizeTilingNoCutTail 首轮 sizeTaken<=MIN_PER_UB_SIZE(1024) 触发 break。 | ||
| 557 | +// 输入 x{2,12,12} FLOAT,diag{2,3,12},k{-1,1}。 | ||
| 558 | +// 期望:成功;tilingKey=0x102(SCATTER+VL满载);tilingData 为 MSDV2NoCutTailTilingData: | ||
| 559 | +// input.coreNum=1, mergeDimSize=2, xRowNum=12, xColNum=12, diagNum=3, maxDiagLen=12, | ||
| 560 | +// k0=-1, k1=1, mergeDimNumPerCore=1, ubFactor=2;blockNum=1;workspace=0。 | ||
| 561 | +TEST_F(MatrixSetDiagV2TilingTest, test_tiling_nocuttail_optimize_break_min_ub) | ||
| 562 | +{ | ||
| 563 | + std::vector<int32_t> kValues = {-1, 1}; | ||
| 564 | + gert::TilingContextPara tilingContextPara("MatrixSetDiagV2", | ||
| 565 | + { | ||
| 566 | + {{{2, 12, 12}, {2, 12, 12}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 567 | + {{{2, 3, 12}, {2, 3, 12}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 568 | + {{{2}, {2}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 569 | + }, | ||
| 570 | + { | ||
| 571 | + {{{2, 12, 12}, {2, 12, 12}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 572 | + }, | ||
| 573 | + &compileInfo); | ||
| 574 | + TilingInfo tilingInfo; | ||
| 575 | + auto tilingRet = ExecuteTiling(tilingContextPara, tilingInfo); | ||
| 576 | + ASSERT_TRUE(tilingRet); | ||
| 577 | + std::vector<int64_t> expectWorkspaces = {0}; | ||
| 578 | + EXPECT_EQ(tilingInfo.workspaceSizes, expectWorkspaces); | ||
| 579 | + EXPECT_EQ(tilingInfo.blockNum, 1); | ||
| 580 | + EXPECT_EQ(tilingInfo.tilingKey, 0b0'0'1'00000010); | ||
| 581 | + ASSERT_EQ(tilingInfo.tilingDataSize, sizeof(MSDV2NoCutTailTilingData)); | ||
| 582 | + MSDV2NoCutTailTilingData* data = reinterpret_cast<MSDV2NoCutTailTilingData*>(tilingInfo.tilingData.get()); | ||
| 583 | + EXPECT_EQ(data->input.coreNum, 1); | ||
| 584 | + EXPECT_EQ(data->input.mergeDimSize, 2); | ||
| 585 | + EXPECT_EQ(data->input.xRowNum, 12); | ||
| 586 | + EXPECT_EQ(data->input.xColNum, 12); | ||
| 587 | + EXPECT_EQ(data->input.diagNum, 3); | ||
| 588 | + EXPECT_EQ(data->input.maxDiagLen, 12); | ||
| 589 | + EXPECT_EQ(data->input.k0, -1); | ||
| 590 | + EXPECT_EQ(data->input.k1, 1); | ||
| 591 | + EXPECT_EQ(data->mergeDimNumPerCore, 1); | ||
| 592 | + EXPECT_EQ(data->ubFactor, 2); | ||
| 593 | +} | ||
| 594 | + | ||
| 595 | +// 场景:ratio=0.005<SIMT_RATIO 使 way=SIMT 且 ubSize 正常,直接走 | ||
| 596 | +// CalculateValidBufSize SIMT 分支并 FillNoCutTailTilingData(realCoreNum=64 不进优化)。 | ||
| 597 | +// 输入 x{64,200,100} FLOAT,diag{64,100},k{0}。 | ||
| 598 | +// 期望:成功;tilingKey=0x003(SIMT);tilingData 为 MSDV2NoCutTailTilingData: | ||
| 599 | +// input.coreNum=64, mergeDimSize=64, xRowNum=200, xColNum=100, diagNum=1, maxDiagLen=100, | ||
| 600 | +// k0=0, k1=0, mergeDimNumPerCore=1, ubFactor=1;blockNum=64;workspace=0。 | ||
| 601 | +TEST_F(MatrixSetDiagV2TilingTest, test_tiling_simt_direct_nocuttail) | ||
| 602 | +{ | ||
| 603 | + std::vector<int32_t> kValues = {0}; | ||
| 604 | + gert::TilingContextPara tilingContextPara("MatrixSetDiagV2", | ||
| 605 | + { | ||
| 606 | + {{{64, 200, 100}, {64, 200, 100}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 607 | + {{{64, 100}, {64, 100}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 608 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 609 | + }, | ||
| 610 | + { | ||
| 611 | + {{{64, 200, 100}, {64, 200, 100}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 612 | + }, | ||
| 613 | + &compileInfo); | ||
| 614 | + TilingInfo tilingInfo; | ||
| 615 | + auto tilingRet = ExecuteTiling(tilingContextPara, tilingInfo); | ||
| 616 | + ASSERT_TRUE(tilingRet); | ||
| 617 | + std::vector<int64_t> expectWorkspaces = {0}; | ||
| 618 | + EXPECT_EQ(tilingInfo.workspaceSizes, expectWorkspaces); | ||
| 619 | + EXPECT_EQ(tilingInfo.blockNum, 64); | ||
| 620 | + EXPECT_EQ(tilingInfo.tilingKey, 0b0'0'0'00000011); | ||
| 621 | + ASSERT_EQ(tilingInfo.tilingDataSize, sizeof(MSDV2NoCutTailTilingData)); | ||
| 622 | + MSDV2NoCutTailTilingData* data = reinterpret_cast<MSDV2NoCutTailTilingData*>(tilingInfo.tilingData.get()); | ||
| 623 | + EXPECT_EQ(data->input.coreNum, 64); | ||
| 624 | + EXPECT_EQ(data->input.mergeDimSize, 64); | ||
| 625 | + EXPECT_EQ(data->input.xRowNum, 200); | ||
| 626 | + EXPECT_EQ(data->input.xColNum, 100); | ||
| 627 | + EXPECT_EQ(data->input.diagNum, 1); | ||
| 628 | + EXPECT_EQ(data->input.maxDiagLen, 100); | ||
| 629 | + EXPECT_EQ(data->input.k0, 0); | ||
| 630 | + EXPECT_EQ(data->input.k1, 0); | ||
| 631 | + EXPECT_EQ(data->mergeDimNumPerCore, 1); | ||
| 632 | + EXPECT_EQ(data->ubFactor, 1); | ||
| 633 | +} | ||
| @@ -204,3 +204,135 @@ TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_test11) | |||
| 204 | }; | 204 | }; |
| 205 | ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | 205 | ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); |
| 206 | } | 206 | } |
| 207 | + | ||
| 208 | +// 场景:k 为非常量张量(不携带数据),diagonal 维度(1) 小于 input 维度(3) - 1,即 3-1=2 > 1。 | ||
| 209 | +// 期望:CheckShape 中非 const k 分支的 xDimNum_-1 > diagDimNum_ 校验失败,返回 GRAPH_FAILED。 | ||
| 210 | +TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_nonconst_k_diag_dim_too_small) | ||
| 211 | +{ | ||
| 212 | + gert::InfershapeContextPara infershapeContextPara("MatrixSetDiagV2", | ||
| 213 | + { | ||
| 214 | + {{{3, 4, 4}, {3, 4, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 215 | + {{{4}, {4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 216 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND}, | ||
| 217 | + }, | ||
| 218 | + { | ||
| 219 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 220 | + }); | ||
| 221 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED); | ||
| 222 | +} | ||
| 223 | + | ||
| 224 | +// 场景:k 为非常量张量(不携带数据),diagonal 维度(4) 大于 input 维度(3)。 | ||
| 225 | +// 期望:CheckShape 中非 const k 分支的 diagDimNum_ > xDimNum_ 校验失败,返回 GRAPH_FAILED。 | ||
| 226 | +TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_nonconst_k_diag_dim_too_large) | ||
| 227 | +{ | ||
| 228 | + gert::InfershapeContextPara infershapeContextPara("MatrixSetDiagV2", | ||
| 229 | + { | ||
| 230 | + {{{3, 4, 4}, {3, 4, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 231 | + {{{3, 4, 4, 4}, {3, 4, 4, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 232 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND}, | ||
| 233 | + }, | ||
| 234 | + { | ||
| 235 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 236 | + }); | ||
| 237 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED); | ||
| 238 | +} | ||
| 239 | + | ||
| 240 | +// 场景:k 为非常量张量(不携带数据),diagonal 维度(2) 介于 input 维度(3)-1 与 input 维度(3) 之间, | ||
| 241 | +// 且 diagonal 前导维(3) 与 input 前导维(3) 相等。 | ||
| 242 | +// 期望:CheckShape 非 const k 分支校验通过并成功返回,输出 shape 为 {3,4,4}。 | ||
| 243 | +TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_nonconst_k_success) | ||
| 244 | +{ | ||
| 245 | + gert::InfershapeContextPara infershapeContextPara("MatrixSetDiagV2", | ||
| 246 | + { | ||
| 247 | + {{{3, 4, 4}, {3, 4, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 248 | + {{{3, 4}, {3, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 249 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND}, | ||
| 250 | + }, | ||
| 251 | + { | ||
| 252 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 253 | + }); | ||
| 254 | + std::vector<std::vector<int64_t>> expectOutputShape = { | ||
| 255 | + {3, 4, 4}, | ||
| 256 | + }; | ||
| 257 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 258 | +} | ||
| 259 | + | ||
| 260 | +// 场景:k 为常量 {0} 且 const 张量校验通过;diagonal 前导维为 UNKNOWN_DIM(-1),其余维度合法。 | ||
| 261 | +// 期望:SetOutputShape 循环中遇到 diag 未知前导维时 continue,输出 shape 保持 {3,4,4}。 | ||
| 262 | +TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_diag_leading_dim_unknown) | ||
| 263 | +{ | ||
| 264 | + std::vector<int32_t> kValues = {0}; | ||
| 265 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 266 | + "MatrixSetDiagV2", | ||
| 267 | + { | ||
| 268 | + {{{3, 4, 4}, {3, 4, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 269 | + {{{-1, 4}, {-1, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 270 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 271 | + }, | ||
| 272 | + { | ||
| 273 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 274 | + }); | ||
| 275 | + std::vector<std::vector<int64_t>> expectOutputShape = { | ||
| 276 | + {3, 4, 4}, | ||
| 277 | + }; | ||
| 278 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 279 | +} | ||
| 280 | + | ||
| 281 | +// 场景:k 为常量 {0} 且 const 张量校验通过;input 前导维为 UNKNOWN_DIM(-1),diagonal 前导维已知为 3。 | ||
| 282 | +// 期望:SetOutputShape 循环中当 x 前导维未知时用 diagonal 前导维填充,输出 shape 为 {3,4,4}。 | ||
| 283 | +TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_x_leading_dim_unknown_fill_diag) | ||
| 284 | +{ | ||
| 285 | + std::vector<int32_t> kValues = {0}; | ||
| 286 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 287 | + "MatrixSetDiagV2", | ||
| 288 | + { | ||
| 289 | + {{{-1, 4, 4}, {-1, 4, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 290 | + {{{3, 4}, {3, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 291 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND, true, kValues.data()}, | ||
| 292 | + }, | ||
| 293 | + { | ||
| 294 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 295 | + }); | ||
| 296 | + std::vector<std::vector<int64_t>> expectOutputShape = { | ||
| 297 | + {3, 4, 4}, | ||
| 298 | + }; | ||
| 299 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 300 | +} | ||
| 301 | + | ||
| 302 | +// 场景:input 为未知 rank({-2}),diagonal/k 正常。 | ||
| 303 | +// 期望:Inference 检测到 x 未知 rank 后直接 SetUnknownRank(y),成功返回,输出 shape 为 {-2}。 | ||
| 304 | +TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_x_unknown_rank) | ||
| 305 | +{ | ||
| 306 | + gert::InfershapeContextPara infershapeContextPara("MatrixSetDiagV2", | ||
| 307 | + { | ||
| 308 | + {{{-2}, {-2}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 309 | + {{{3, 4}, {3, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 310 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND}, | ||
| 311 | + }, | ||
| 312 | + { | ||
| 313 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 314 | + }); | ||
| 315 | + std::vector<std::vector<int64_t>> expectOutputShape = { | ||
| 316 | + {-2}, | ||
| 317 | + }; | ||
| 318 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 319 | +} | ||
| 320 | + | ||
| 321 | +// 场景:diagonal 为未知 rank({-2}),input/k 正常。 | ||
| 322 | +// 期望:Inference 检测到 diag 未知 rank 后直接 SetUnknownRank(y),成功返回,输出 shape 为 {-2}。 | ||
| 323 | +TEST_F(MatrixSetDiagV2Infershape, matrix_set_diag_infershape_diag_unknown_rank) | ||
| 324 | +{ | ||
| 325 | + gert::InfershapeContextPara infershapeContextPara("MatrixSetDiagV2", | ||
| 326 | + { | ||
| 327 | + {{{3, 4, 4}, {3, 4, 4}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 328 | + {{{-2}, {-2}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 329 | + {{{1}, {1}}, ge::DT_INT32, ge::FORMAT_ND}, | ||
| 330 | + }, | ||
| 331 | + { | ||
| 332 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 333 | + }); | ||
| 334 | + std::vector<std::vector<int64_t>> expectOutputShape = { | ||
| 335 | + {-2}, | ||
| 336 | + }; | ||
| 337 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 338 | +} | ||