已合并
feat: add matrix_set_diag_v2 UT #4608
zhanw_coding创建于 8月13日
feat: add matrix_set_diag_v2 UT #4608
已合并
zhanw_coding创建于 8月13日
共 5 个文件变更+423-2
@@ -938,6 +938,9 @@ VC5@ops-math:
938 # depth_to_space938 # 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_pad944 # 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/*.cpp1013 - ops/ops-math/conversion/depth_to_space/examples/*.cpp
1011 - ops/ops-math/conversion/depth_to_space/tests/**/*.cpp1014 - ops/ops-math/conversion/depth_to_space/tests/**/*.cpp
1012 - ops/ops-math/conversion/depth_to_space/tests/**/*.py1015 - 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_pad1019 # 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.cpp1104 - 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.cpp1110 - 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+ 
155TEST_F(MatrixSetDiagTilingTest, test_tiling_failed_batch_dim_invalid)174TEST_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+}