已合并
修复自定义算子包安装失败问题 #9026
修复自定义算子包安装失败问题 #9026
已合并
chenyifan创建于 7月22日
7 个文件变更+93-47
@@ -591,6 +591,10 @@ function(add_tiling_modules)
591 c_sec591 c_sec
592 )592 )
593 endif()593 endif()
594+ # Supplement the existing tiling_obj with built-in package macros for C++ conditional compilation.
595+ if(ENABLE_BUILT_IN AND TARGET ${OPHOST_NAME}_tiling_obj)
596+ target_compile_definitions(${OPHOST_NAME}_tiling_obj PRIVATE ENABLE_BUILT_IN=1)
597+ endif()
594endfunction()598endfunction()
595 599 
596function(add_graph_plugin_modules)600function(add_graph_plugin_modules)
@@ -52,7 +52,7 @@ TEST_F(AllGatherMatmulV2InferShapeTest, Basic)
52 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},52 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
53 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},53 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
54 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},54 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
55- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},55+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
56 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},56 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
57 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},57 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
58 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},58 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
@@ -91,7 +91,7 @@ TEST_F(AllGatherMatmulV2InferShapeTest, EmptyTensorFailTest)
91 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},91 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
92 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},92 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
93 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},93 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
94- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},94+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
95 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},95 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
96 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},96 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
97 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},97 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
@@ -134,7 +134,7 @@ TEST_F(AllGatherMatmulV2InferShapeTest, Pertensor)
134 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},134 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
135 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},135 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
136 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},136 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
137- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},137+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
138 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},138 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
139 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},139 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
140 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},140 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
@@ -178,7 +178,7 @@ TEST_F(AllGatherMatmulV2InferShapeTest, Perblock)
178 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},178 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
179 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},179 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
180 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},180 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
181- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},181+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
182 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},182 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
183 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},183 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
184 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},184 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
@@ -264,7 +264,7 @@ TEST_F(AllGatherMatmulV2InferDTypeTest, AttrYDtype)
264 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},264 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
265 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},265 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
266 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},266 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
267- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},267+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
268 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},268 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
269 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},269 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
270 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},270 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
@@ -292,7 +292,7 @@ TEST_F(AllGatherMatmulV2InferShapeTest, AmaxOutTrue)
292 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},292 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
293 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},293 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
294 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},294 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
295- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},295+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
296 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},296 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
297 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},297 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
298 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},298 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
@@ -327,7 +327,7 @@ TEST_F(AllGatherMatmulV2InferShapeTest, AmaxOutTrueFP8)
327 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},327 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
328 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},328 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
329 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},329 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
330- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},330+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
331 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},331 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
332 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},332 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
333 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},333 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
@@ -356,7 +356,7 @@ TEST_F(AllGatherMatmulV2InferDTypeTest, FP8E4M3FN)
356 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},356 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
357 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},357 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
358 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},358 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
359- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},359+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
360 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},360 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
361 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},361 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
362 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},362 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
@@ -387,7 +387,7 @@ TEST_F(AllGatherMatmulV2InferDTypeTest, HIFLOAT8)
387 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},387 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
388 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},388 {"gather_index", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
389 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},389 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
390- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},390+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
391 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},391 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
392 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},392 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
393 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},393 {"is_gather_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},
@@ -17,7 +17,6 @@ if (BUILD_OPEN_PROJECT)
17 op_host/op_tiling/matmul_formulaic_tiling.cpp17 op_host/op_tiling/matmul_formulaic_tiling.cpp
18 op_host/op_tiling/matmul_performance.cpp18 op_host/op_tiling/matmul_performance.cpp
19 op_host/op_tiling/mc2_fit_based_balance_tiling.cpp19 op_host/op_tiling/mc2_fit_based_balance_tiling.cpp
20- op_host/mc2_common_infershape.cpp
21 utils/mc2_hcom_topo_info.cpp20 utils/mc2_hcom_topo_info.cpp
22 utils/mc2_log.cpp21 utils/mc2_log.cpp
23 op_host/op_tiling/mc2_matmul_tiling_cfg.cpp22 op_host/op_tiling/mc2_matmul_tiling_cfg.cpp
@@ -27,6 +26,10 @@ if (BUILD_OPEN_PROJECT)
27 op_host/op_tiling/one_calc_two_comm_tiling.cpp26 op_host/op_tiling/one_calc_two_comm_tiling.cpp
28 )27 )
29 28 
29+ if(ENABLE_BUILT_IN)
30+ list(APPEND OPTILING_SRCS op_host/mc2_common_infershape.cpp)
31+ endif()
32+ 
30 if(UT_TEST_ALL OR OP_HOST_UT)33 if(UT_TEST_ALL OR OP_HOST_UT)
31 list(FILTER OPTILING_SRCS EXCLUDE REGEX "utils/mc2_hcom_topo_info\\.cpp")34 list(FILTER OPTILING_SRCS EXCLUDE REGEX "utils/mc2_hcom_topo_info\\.cpp")
32 file(GLOB MC2_TILING_UT_STUB ${UT_COMMON_INC}/mc2_hcom_topology_mocker.cpp ${UT_COMMON_INC}/op_cache_tiling.cpp ${UT_COMMON_INC}/rt_soc_spec_mocker.cpp)35 file(GLOB MC2_TILING_UT_STUB ${UT_COMMON_INC}/mc2_hcom_topology_mocker.cpp ${UT_COMMON_INC}/op_cache_tiling.cpp ${UT_COMMON_INC}/rt_soc_spec_mocker.cpp)
@@ -45,6 +48,15 @@ if (BUILD_OPEN_PROJECT)
45 target_link_directories(${OPHOST_NAME}_tiling_obj PRIVATE mc2_tiling_ut_stub)48 target_link_directories(${OPHOST_NAME}_tiling_obj PRIVATE mc2_tiling_ut_stub)
46 endif()49 endif()
47 endif()50 endif()
51+ if(NOT ENABLE_BUILT_IN)
52+ set(MC2_INFER_SRCS
53+ op_host/mc2_common_infershape.cpp
54+ )
55+ if(MC2_INFER_SRCS)
56+ add_infer_modules()
57+ target_sources(${OPHOST_NAME}_infer_obj PRIVATE ${MC2_INFER_SRCS})
58+ endif()
59+ endif()
48 file(GLOB COMMON_UTILS_SRCS60 file(GLOB COMMON_UTILS_SRCS
49 utils/mc2_platform_info.cpp61 utils/mc2_platform_info.cpp
50 )62 )
@@ -18,9 +18,7 @@
18 18 
19using namespace ge;19using namespace ge;
20namespace ops {20namespace ops {
21-// infershape 公共函数21+static ge::graphStatus CheckMatrixInputShapes(const gert::InferShapeContext* context, CommParas& commParas)
22-ge::graphStatus CommonParamCheck(
23- const gert::InferShapeContext* context, const size_t isTransAIndex, const size_t isTransBIndex, CommParas& commParas)
24{22{
25 commParas.x1MatrixShape = context->GetInputShape(0);23 commParas.x1MatrixShape = context->GetInputShape(0);
26 OPS_CHECK_NULL_WITH_CONTEXT(context, commParas.x1MatrixShape);24 OPS_CHECK_NULL_WITH_CONTEXT(context, commParas.x1MatrixShape);
@@ -36,6 +34,56 @@ ge::graphStatus CommonParamCheck(
36 std::to_string(commParas.x2MatrixShape->GetDimNum()) + "D", "2D");34 std::to_string(commParas.x2MatrixShape->GetDimNum()) + "D", "2D");
37 return ge::GRAPH_FAILED;35 return ge::GRAPH_FAILED;
38 }36 }
37+ return ge::GRAPH_SUCCESS;
38+}
39+ 
40+static ge::graphStatus ResolveRankSize(
41+ const gert::InferShapeContext* context, const char* groupStr, const int64_t* rankSizeAttr, int64_t& rankSize)
42+{
43+ if (*rankSizeAttr <= 0) {
44+#if defined(ENABLE_BUILT_IN)
45+ uint32_t rankNum = 0;
46+ if (Mc2Hcom::MC2HcomTopology::CommGetInstSizeByGroup(groupStr, &rankNum) != HCCL_SUCCESS || rankNum == 0) {
47+ OP_LOGE(context->GetNodeName(), "Get rank size failed, group [%s], rankSize [%u]", groupStr, rankNum);
48+ return ge::GRAPH_FAILED;
49+ }
50+ rankSize = static_cast<int64_t>(rankNum);
51+#else
52+ OP_LOGE(context->GetNodeName(), "Get rank size failed, rankSize [%lld]", *rankSizeAttr);
53+ return ge::GRAPH_FAILED;
54+#endif
55+ } else {
56+ rankSize = *rankSizeAttr;
57+ }
58+ return ge::GRAPH_SUCCESS;
59+}
60+ 
61+static void FillMatmulDims(CommParas& commParas, bool isTransA, bool isTransB)
62+{
63+ commParas.dimM = !isTransA ? commParas.x1MatrixShape->GetDim(0) : commParas.x1MatrixShape->GetDim(1);
64+ commParas.dimKX1 = !isTransA ? commParas.x1MatrixShape->GetDim(1) : commParas.x1MatrixShape->GetDim(0);
65+ commParas.dimKX2 = !isTransB ? commParas.x2MatrixShape->GetDim(0) : commParas.x2MatrixShape->GetDim(1);
66+ commParas.dimN = !isTransB ? commParas.x2MatrixShape->GetDim(1) : commParas.x2MatrixShape->GetDim(0);
67+}
68+ 
69+static ge::graphStatus CheckKDimMatch(const gert::InferShapeContext* context, const CommParas& commParas)
70+{
71+ if (commParas.dimKX1 != commParas.dimKX2) {
72+ OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(context->GetNodeName(), "x1.k, x2.k",
73+ std::to_string(commParas.dimKX1) + ", " + std::to_string(commParas.dimKX2),
74+ "The values of x1.k and x2.k must be the same");
75+ return ge::GRAPH_FAILED;
76+ }
77+ return ge::GRAPH_SUCCESS;
78+}
79+ 
80+// infershape 公共函数
81+ge::graphStatus CommonParamCheck(
82+ const gert::InferShapeContext* context, const size_t isTransAIndex, const size_t isTransBIndex, CommParas& commParas)
83+{
84+ if (CheckMatrixInputShapes(context, commParas) != ge::GRAPH_SUCCESS) {
85+ return ge::GRAPH_FAILED;
86+ }
39 auto attrs = context->GetAttrs();87 auto attrs = context->GetAttrs();
40 OPS_CHECK_NULL_WITH_CONTEXT(context, attrs);88 OPS_CHECK_NULL_WITH_CONTEXT(context, attrs);
41 const bool* isTransA = attrs->GetAttrPointer<bool>(isTransAIndex);89 const bool* isTransA = attrs->GetAttrPointer<bool>(isTransAIndex);
@@ -43,41 +91,24 @@ ge::graphStatus CommonParamCheck(
43 const int64_t* rankSizeAttr = attrs->GetAttrPointer<int64_t>(RANK_SIZE);91 const int64_t* rankSizeAttr = attrs->GetAttrPointer<int64_t>(RANK_SIZE);
44 92 
45 const char* groupStr = attrs->GetAttrPointer<char>(GROUP);93 const char* groupStr = attrs->GetAttrPointer<char>(GROUP);
94+ 
46 if (groupStr == nullptr) {95 if (groupStr == nullptr) {
47 OP_LOGE_WITH_INVALID_INPUT(context->GetNodeName(), "groupStr");96 OP_LOGE_WITH_INVALID_INPUT(context->GetNodeName(), "groupStr");
48 return ge::GRAPH_FAILED;97 return ge::GRAPH_FAILED;
49 }98 }
50- commParas.rankSize = -1;99+ 
51- uint32_t rankNum = 0;100+ if (ResolveRankSize(context, groupStr, rankSizeAttr, commParas.rankSize) != ge::GRAPH_SUCCESS) {
52- if (*rankSizeAttr <= 0) {101+ return ge::GRAPH_FAILED;
53- if ((Mc2Hcom::MC2HcomTopology::CommGetInstSizeByGroup(groupStr, &rankNum)) != HCCL_SUCCESS || rankNum == 0) {
54- OP_LOGE(context->GetNodeName(), "Get rank size failed, group [%s], rankSize [%u]", groupStr, rankNum);
55- return ge::GRAPH_FAILED;
56- } else {
57- commParas.rankSize = rankNum;
58- }
59- commParas.rankSize = static_cast<int64_t>(rankNum);
60- } else {
61- commParas.rankSize = *rankSizeAttr;
62 }102 }
63 103 
64- commParas.dimM = !(*isTransA) ? commParas.x1MatrixShape->GetDim(0) : commParas.x1MatrixShape->GetDim(1);104+ FillMatmulDims(commParas, *isTransA, *isTransB);
65- commParas.dimKX1 = !(*isTransA) ? commParas.x1MatrixShape->GetDim(1) : commParas.x1MatrixShape->GetDim(0);
66- commParas.dimKX2 = !(*isTransB) ? commParas.x2MatrixShape->GetDim(0) : commParas.x2MatrixShape->GetDim(1);
67- commParas.dimN = !(*isTransB) ? commParas.x2MatrixShape->GetDim(1) : commParas.x2MatrixShape->GetDim(0);
68- 
69 OP_LOGI(105 OP_LOGI(
70 context->GetNodeName(),106 context->GetNodeName(),
71 "group = %s isTransA %d isTransB %d x1.M = [%ld] x1.K = [%ld]"107 "group = %s isTransA %d isTransB %d x1.M = [%ld] x1.K = [%ld]"
72 " x2.K = [%ld] x2.N = [%ld] rankSize = [%ld].",108 " x2.K = [%ld] x2.N = [%ld] rankSize = [%ld].",
73 groupStr, (*isTransA), (*isTransB), commParas.x1MatrixShape->GetDim(0), commParas.x1MatrixShape->GetDim(1),109 groupStr, (*isTransA), (*isTransB), commParas.x1MatrixShape->GetDim(0), commParas.x1MatrixShape->GetDim(1),
74 commParas.x2MatrixShape->GetDim(0), commParas.x2MatrixShape->GetDim(1), commParas.rankSize);110 commParas.x2MatrixShape->GetDim(0), commParas.x2MatrixShape->GetDim(1), commParas.rankSize);
75- if (commParas.dimKX1 != commParas.dimKX2) {111+ return CheckKDimMatch(context, commParas);
76- OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(context->GetNodeName(), "x1.k, x2.k",
77- std::to_string(commParas.dimKX1) + ", " + std::to_string(commParas.dimKX2), "The values of x1.k and x2.k must be the same");
78- return ge::GRAPH_FAILED;
79- }
80- return ge::GRAPH_SUCCESS;
81}112}
82 113 
83ge::graphStatus AllGatherMatmulInferYShape(gert::InferShapeContext* context, CommParas& commParas)114ge::graphStatus AllGatherMatmulInferYShape(gert::InferShapeContext* context, CommParas& commParas)
@@ -85,7 +116,6 @@ ge::graphStatus AllGatherMatmulInferYShape(gert::InferShapeContext* context, Com
85 OP_LOGE_IF(116 OP_LOGE_IF(
86 CommonParamCheck(context, AG_IS_TRANS_A, AG_IS_TRANS_B, commParas) != GRAPH_SUCCESS, GRAPH_FAILED,117 CommonParamCheck(context, AG_IS_TRANS_A, AG_IS_TRANS_B, commParas) != GRAPH_SUCCESS, GRAPH_FAILED,
87 context->GetNodeName(), "CommonParamCheck excute failed.");118 context->GetNodeName(), "CommonParamCheck excute failed.");
88- // 动态shape入图时 m轴-1时,不再进行(dimM * rankSize)的处理
89 if (commParas.dimM == -1) {119 if (commParas.dimM == -1) {
90 commParas.rankSize = 1;120 commParas.rankSize = 1;
91 }121 }
@@ -145,7 +175,6 @@ ge::graphStatus InferMatmulReduceScatterCommon(gert::InferShapeContext* context)
145 OP_LOGE_IF(175 OP_LOGE_IF(
146 CommonParamCheck(context, RS_IS_TRANS_A, RS_IS_TRANS_B, commParas) != GRAPH_SUCCESS, GRAPH_FAILED,176 CommonParamCheck(context, RS_IS_TRANS_A, RS_IS_TRANS_B, commParas) != GRAPH_SUCCESS, GRAPH_FAILED,
147 context->GetNodeName(), "CommonParamCheck excute failed.");177 context->GetNodeName(), "CommonParamCheck excute failed.");
148- // 动态shape入图时 m轴-1时,不再进行(dimM / rankSize)的处理
149 if (commParas.dimM == -1) {178 if (commParas.dimM == -1) {
150 commParas.rankSize = 1;179 commParas.rankSize = 1;
151 }180 }
@@ -161,4 +190,4 @@ ge::graphStatus InferMatmulReduceScatterCommon(gert::InferShapeContext* context)
161 yShape->SetDim(1, commParas.dimN);190 yShape->SetDim(1, commParas.dimN);
162 return GRAPH_SUCCESS;191 return GRAPH_SUCCESS;
163}192}
164-} // namespace ops193+} // namespace ops
@@ -18,8 +18,9 @@
18 18 
19#include "mc2_log.h"19#include "mc2_log.h"
20#include "register/op_impl_registry.h"20#include "register/op_impl_registry.h"
21+#if defined(ENABLE_BUILT_IN)
21#include "mc2_hcom_topo_info.h"22#include "mc2_hcom_topo_info.h"
22- 23+#endif
23namespace ops {24namespace ops {
24 const size_t GROUP = 0;25 const size_t GROUP = 0;
25 const size_t AG_IS_TRANS_A = 1;26 const size_t AG_IS_TRANS_A = 1;
@@ -48,7 +48,7 @@ TEST_F(MatmulReduceScatterInfershape, Basic)
48 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},48 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
49 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},49 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
50 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},50 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
51- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)}51+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)}
52 }52 }
53 );53 );
54 Mc2Hcom::MockValues hcomTopologyMockValues {54 Mc2Hcom::MockValues hcomTopologyMockValues {
@@ -49,7 +49,7 @@ TEST_F(MatmulReduceScatterV2InferShapeTest, Basic)
49 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},49 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
50 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},50 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
51 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},51 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
52- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},52+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
53 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},53 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
54 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},54 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
55 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},55 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
@@ -90,7 +90,7 @@ TEST_F(MatmulReduceScatterV2InferShapeTest, EmptyTensorTest)
90 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},90 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
91 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},91 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
92 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},92 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
93- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},93+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
94 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},94 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
95 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},95 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
96 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},96 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
@@ -133,7 +133,7 @@ TEST_F(MatmulReduceScatterV2InferShapeTest, Pertensor)
133 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},133 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
134 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},134 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
135 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},135 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
136- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},136+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
137 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},137 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
138 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},138 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
139 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},139 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
@@ -176,7 +176,7 @@ TEST_F(MatmulReduceScatterV2InferShapeTest, Perblock)
176 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},176 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
177 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},177 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
178 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},178 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
179- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},179+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
180 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},180 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
181 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},181 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
182 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},182 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
@@ -250,7 +250,7 @@ TEST_F(MatmulReduceScatterV2InferDTypeTest, AttrYDtype)
250 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},250 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
251 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},251 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
252 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},252 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
253- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},253+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
254 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},254 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
255 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},255 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
256 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},256 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
@@ -290,7 +290,7 @@ TEST_F(MatmulReduceScatterV2InferShapeTest, AmaxOutEnabled)
290 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},290 {"is_trans_a", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
291 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},291 {"is_trans_b", Ops::Transformer::AnyValue::CreateFrom<bool>(false)},
292 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},292 {"comm_turn", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
293- {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},293+ {"rank_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(8)},
294 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},294 {"block_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
295 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},295 {"group_size", Ops::Transformer::AnyValue::CreateFrom<int64_t>(0)},
296 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},296 {"is_amax_out", Ops::Transformer::AnyValue::CreateFrom<bool>(true)},