已合并
修复自定义算子包安装失败问题 #9026
chenyifan创建于 7月22日
修复自定义算子包安装失败问题 #9026
已合并
共 7 个文件变更+93-47
| @@ -591,6 +591,10 @@ function(add_tiling_modules) | |||
| 591 | c_sec | 591 | 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() | ||
| 594 | endfunction() | 598 | endfunction() |
| 595 | 599 | ||
| 596 | function(add_graph_plugin_modules) | 600 | function(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.cpp | 17 | op_host/op_tiling/matmul_formulaic_tiling.cpp |
| 18 | op_host/op_tiling/matmul_performance.cpp | 18 | op_host/op_tiling/matmul_performance.cpp |
| 19 | op_host/op_tiling/mc2_fit_based_balance_tiling.cpp | 19 | op_host/op_tiling/mc2_fit_based_balance_tiling.cpp |
| 20 | - op_host/mc2_common_infershape.cpp | ||
| 21 | utils/mc2_hcom_topo_info.cpp | 20 | utils/mc2_hcom_topo_info.cpp |
| 22 | utils/mc2_log.cpp | 21 | utils/mc2_log.cpp |
| 23 | op_host/op_tiling/mc2_matmul_tiling_cfg.cpp | 22 | 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.cpp | 26 | 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_SRCS | 60 | file(GLOB COMMON_UTILS_SRCS |
| 49 | utils/mc2_platform_info.cpp | 61 | utils/mc2_platform_info.cpp |
| 50 | ) | 62 | ) |
| @@ -18,9 +18,7 @@ | |||
| 18 | 18 | ||
| 19 | using namespace ge; | 19 | using namespace ge; |
| 20 | namespace ops { | 20 | namespace 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 | + | ||
| 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 | + | ||
| 52 | + OP_LOGE(context->GetNodeName(), "Get rank size failed, rankSize [%lld]", *rankSizeAttr); | ||
| 53 | + return ge::GRAPH_FAILED; | ||
| 54 | + | ||
| 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 | ||
| 83 | ge::graphStatus AllGatherMatmulInferYShape(gert::InferShapeContext* context, CommParas& commParas) | 114 | ge::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 ops | 193 | +} // namespace ops |
| @@ -18,8 +18,9 @@ | |||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | + | ||
| 21 | 22 | ||
| 22 | - | 23 | +#endif |
| 23 | namespace ops { | 24 | namespace 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)}, |