已合并
性能建模回退 #2052
高煜博创建于 9月11日
性能建模回退 #2052
已合并
共 23 个文件变更+24-4633
| @@ -41,7 +41,6 @@ struct NodeDetail { | |||
| 41 | Expr gm_stride{CreateExpr(0)}; | 41 | Expr gm_stride{CreateExpr(0)}; |
| 42 | Expr ub_stride{CreateExpr(0)}; | 42 | Expr ub_stride{CreateExpr(0)}; |
| 43 | ascir_param::CastNodeParams cast_node_params; | 43 | ascir_param::CastNodeParams cast_node_params; |
| 44 | - ascir_param::BroadcastNodeParams broadcast_node_params; | ||
| 45 | ascir_param::CompareNodeParams compare_node_params; | 44 | ascir_param::CompareNodeParams compare_node_params; |
| 46 | ascir_param::WhereNodeParams where_node_params; | 45 | ascir_param::WhereNodeParams where_node_params; |
| 47 | ascir_param::UnaryBitWidthChangeNodeParams unary_bitwidth_change_node_params; | 46 | ascir_param::UnaryBitWidthChangeNodeParams unary_bitwidth_change_node_params; |
| @@ -38,13 +38,6 @@ af::Status FillCastParams(const ascir_param::AscirNodeParams ¶ms, NodeInfo & | |||
| 38 | return af::SUCCESS; | 38 | return af::SUCCESS; |
| 39 | } | 39 | } |
| 40 | 40 | ||
| 41 | -af::Status FillBroadcastParams(const ascir_param::AscirNodeParams ¶ms, NodeInfo &node_info) { | ||
| 42 | - const auto *broadcast_params = ascir_param::GetSpecificParams<ascir_param::BroadcastNodeParams>(params); | ||
| 43 | - GE_ASSERT_NOTNULL(broadcast_params, "Broadcast specific params is null, node[%s].", node_info.name.c_str()); | ||
| 44 | - node_info.broadcast_node_params = *broadcast_params; | ||
| 45 | - return af::SUCCESS; | ||
| 46 | -} | ||
| 47 | - | ||
| 48 | af::Status FillCompareParams(const ascir_param::AscirNodeParams ¶ms, NodeInfo &node_info) { | 41 | af::Status FillCompareParams(const ascir_param::AscirNodeParams ¶ms, NodeInfo &node_info) { |
| 49 | const auto *compare_params = ascir_param::GetSpecificParams<ascir_param::CompareNodeParams>(params); | 42 | const auto *compare_params = ascir_param::GetSpecificParams<ascir_param::CompareNodeParams>(params); |
| 50 | GE_ASSERT_NOTNULL(compare_params, "Compare specific params is null, node[%s].", node_info.name.c_str()); | 43 | GE_ASSERT_NOTNULL(compare_params, "Compare specific params is null, node[%s].", node_info.name.c_str()); |
| @@ -94,9 +87,6 @@ af::Status FillSpecificParams(const af::AscNodePtr &ge_node, NodeInfo &node_info | |||
| 94 | if (node_info.node_type == kCast) { | 87 | if (node_info.node_type == kCast) { |
| 95 | return FillCastParams(*params, node_info); | 88 | return FillCastParams(*params, node_info); |
| 96 | } | 89 | } |
| 97 | - if (node_info.node_type == kBroadcast) { | ||
| 98 | - return FillBroadcastParams(*params, node_info); | ||
| 99 | - } | ||
| 100 | if (node_info.node_type == kGe || node_info.node_type == kEq || node_info.node_type == kNe || | 90 | if (node_info.node_type == kGe || node_info.node_type == kEq || node_info.node_type == kNe || |
| 101 | node_info.node_type == kGt || node_info.node_type == kLe || node_info.node_type == kLt) { | 91 | node_info.node_type == kGt || node_info.node_type == kLe || node_info.node_type == kLt) { |
| 102 | return FillCompareParams(*params, node_info); | 92 | return FillCompareParams(*params, node_info); |
| @@ -184,7 +184,6 @@ struct NodeInfo { | |||
| 184 | ascir_param::VectorFuncNodeParams vector_func_params; | 184 | ascir_param::VectorFuncNodeParams vector_func_params; |
| 185 | ascir_param::ReduceNodeParams reduce_specific_params; | 185 | ascir_param::ReduceNodeParams reduce_specific_params; |
| 186 | ascir_param::CastNodeParams cast_node_params; | 186 | ascir_param::CastNodeParams cast_node_params; |
| 187 | - ascir_param::BroadcastNodeParams broadcast_node_params; | ||
| 188 | ascir_param::CompareNodeParams compare_node_params; | 187 | ascir_param::CompareNodeParams compare_node_params; |
| 189 | ascir_param::WhereNodeParams where_node_params; | 188 | ascir_param::WhereNodeParams where_node_params; |
| 190 | ascir_param::UnaryBitWidthChangeNodeParams unary_bitwidth_change_node_params; | 189 | ascir_param::UnaryBitWidthChangeNodeParams unary_bitwidth_change_node_params; |
| @@ -81,16 +81,6 @@ struct CastNodeParams { | |||
| 81 | std::vector<ge::Expression> input_strides; | 81 | std::vector<ge::Expression> input_strides; |
| 82 | }; | 82 | }; |
| 83 | 83 | ||
| 84 | -struct BroadcastNodeParams { | ||
| 85 | - bool valid{false}; | ||
| 86 | - bool is_scalar{false}; | ||
| 87 | - int32_t const_rank{-1}; | ||
| 88 | - ge::Expression duplicate_count{ge::Symbol(1U)}; | ||
| 89 | - ParamExprRole duplicate_count_role{ParamExprRole::kSemantic}; | ||
| 90 | - std::vector<ParamExprLeaf> dst_shape; | ||
| 91 | - std::vector<ParamExprLeaf> src_shape; | ||
| 92 | -}; | ||
| 93 | - | ||
| 94 | struct CompareNodeParams { | 84 | struct CompareNodeParams { |
| 95 | bool valid{false}; | 85 | bool valid{false}; |
| 96 | bool is_scalar{false}; | 86 | bool is_scalar{false}; |
| @@ -130,8 +120,8 @@ struct TransposeNodeParams { | |||
| 130 | }; | 120 | }; |
| 131 | 121 | ||
| 132 | using AnySpecificParams = | 122 | using AnySpecificParams = |
| 133 | - std::variant<std::monostate, ReduceNodeParams, VectorFuncNodeParams, CastNodeParams, BroadcastNodeParams, | 123 | + std::variant<std::monostate, ReduceNodeParams, VectorFuncNodeParams, CastNodeParams, CompareNodeParams, |
| 134 | - CompareNodeParams, WhereNodeParams, UnaryBitWidthChangeNodeParams, TransposeNodeParams>; | 124 | + WhereNodeParams, UnaryBitWidthChangeNodeParams, TransposeNodeParams>; |
| 135 | 125 | ||
| 136 | struct AscirNodeParams { | 126 | struct AscirNodeParams { |
| 137 | // 扩展属性载荷版本,用于后续兼容。 | 127 | // 扩展属性载荷版本,用于后续兼容。 |
| @@ -25,7 +25,6 @@ namespace { | |||
| 25 | constexpr const char *kAscirNodeParams = "AscirNodeParams"; | 25 | constexpr const char *kAscirNodeParams = "AscirNodeParams"; |
| 26 | constexpr const char *kVectorFunc = "VectorFunc"; | 26 | constexpr const char *kVectorFunc = "VectorFunc"; |
| 27 | constexpr const char *kCast = "Cast"; | 27 | constexpr const char *kCast = "Cast"; |
| 28 | -constexpr const char *kBroadcast = "Broadcast"; | ||
| 29 | 28 | ||
| 30 | bool IsCompareParamSupported(const std::string &api_name) { | 29 | bool IsCompareParamSupported(const std::string &api_name) { |
| 31 | static const std::set<std::string> kCompareTypes = {"Ge", "Eq", "Ne", "Gt", "Le", "Lt"}; | 30 | static const std::set<std::string> kCompareTypes = {"Ge", "Eq", "Ne", "Gt", "Le", "Lt"}; |
| @@ -124,22 +123,6 @@ af::Status RegisterCastAscirNodeParams(const af::AscNodePtr &node) { | |||
| 124 | return RegisterAscirNodeParams(node, params); | 123 | return RegisterAscirNodeParams(node, params); |
| 125 | } | 124 | } |
| 126 | 125 | ||
| 127 | -af::Status RegisterBroadcastAscirNodeParams(const af::AscNodePtr &node) { | ||
| 128 | - GE_ASSERT_NOTNULL(node); | ||
| 129 | - const auto existing_params = GetAscirNodeParams(node); | ||
| 130 | - if (existing_params != nullptr) { | ||
| 131 | - const auto *broadcast_params = GetSpecificParams<BroadcastNodeParams>(*existing_params); | ||
| 132 | - if (broadcast_params != nullptr && broadcast_params->valid) { | ||
| 133 | - return af::SUCCESS; | ||
| 134 | - } | ||
| 135 | - } | ||
| 136 | - auto params = std::make_shared<AscirNodeParams>(); | ||
| 137 | - params->api_name = node->GetType(); | ||
| 138 | - params->status = ParamBuildStatus::kBuilt; | ||
| 139 | - params->specific_params = BroadcastNodeParams{}; | ||
| 140 | - return RegisterAscirNodeParams(node, params); | ||
| 141 | -} | ||
| 142 | - | ||
| 143 | af::Status RegisterCompareAscirNodeParams(const af::AscNodePtr &node) { | 126 | af::Status RegisterCompareAscirNodeParams(const af::AscNodePtr &node) { |
| 144 | GE_ASSERT_NOTNULL(node); | 127 | GE_ASSERT_NOTNULL(node); |
| 145 | auto params = std::make_shared<AscirNodeParams>(); | 128 | auto params = std::make_shared<AscirNodeParams>(); |
| @@ -566,9 +549,6 @@ af::Status EnrichAscirNodeParams(const AscirParamSourceContext &source) { | |||
| 566 | if (source.node->GetType() == kCast) { | 549 | if (source.node->GetType() == kCast) { |
| 567 | return RegisterCastAscirNodeParams(source.node); | 550 | return RegisterCastAscirNodeParams(source.node); |
| 568 | } | 551 | } |
| 569 | - if (source.node->GetType() == kBroadcast) { | ||
| 570 | - return RegisterBroadcastAscirNodeParams(source.node); | ||
| 571 | - } | ||
| 572 | if (IsCompareParamSupported(source.node->GetType())) { | 552 | if (IsCompareParamSupported(source.node->GetType())) { |
| 573 | return RegisterCompareAscirNodeParams(source.node); | 553 | return RegisterCompareAscirNodeParams(source.node); |
| 574 | } | 554 | } |
| @@ -17,7 +17,6 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | 21 | ||
| 23 | 22 | ||
| @@ -94,38 +93,6 @@ Status BuildReduceAscendGraphND(AscGraph &graph) { | |||
| 94 | } // namespace af | 93 | } // namespace af |
| 95 | namespace att { | 94 | namespace att { |
| 96 | namespace { | 95 | namespace { |
| 97 | - | ||
| 98 | -TEST(SpecificParamsBuilderTest, FillBroadcastParamsCopiesPayloadToNodeInfo) { | ||
| 99 | - af::AscGraph graph("test_graph"); | ||
| 100 | - af::ascir_op::Broadcast broadcast_op("broadcast"); | ||
| 101 | - graph.AddNode(broadcast_op); | ||
| 102 | - auto node = graph.FindNode("broadcast"); | ||
| 103 | - ASSERT_NE(node, nullptr); | ||
| 104 | - | ||
| 105 | - auto params = std::make_shared<ascir_param::AscirNodeParams>(); | ||
| 106 | - params->api_name = "Broadcast"; | ||
| 107 | - params->status = ascir_param::ParamBuildStatus::kBuilt; | ||
| 108 | - ascir_param::BroadcastNodeParams broadcast; | ||
| 109 | - broadcast.valid = true; | ||
| 110 | - broadcast.is_scalar = true; | ||
| 111 | - broadcast.duplicate_count = ge::Symbol(8); | ||
| 112 | - broadcast.src_shape = {{ge::Symbol(1), ascir_param::ParamExprRole::kSemantic}, | ||
| 113 | - {ge::Symbol(8), ascir_param::ParamExprRole::kActualSize}}; | ||
| 114 | - broadcast.dst_shape = {{ge::Symbol(4), ascir_param::ParamExprRole::kActualSize}, | ||
| 115 | - {ge::Symbol(8), ascir_param::ParamExprRole::kActualSize}}; | ||
| 116 | - params->specific_params = broadcast; | ||
| 117 | - ASSERT_TRUE(node->GetOpDesc()->SetExtAttr("AscirNodeParams", params)); | ||
| 118 | - | ||
| 119 | - NodeInfo node_info; | ||
| 120 | - node_info.name = "broadcast"; | ||
| 121 | - node_info.node_type = "Broadcast"; | ||
| 122 | - ASSERT_EQ(FillSpecificParams(node, node_info), af::SUCCESS); | ||
| 123 | - EXPECT_TRUE(node_info.broadcast_node_params.valid); | ||
| 124 | - EXPECT_TRUE(node_info.broadcast_node_params.is_scalar); | ||
| 125 | - EXPECT_EQ(node_info.broadcast_node_params.duplicate_count, ge::Symbol(8)); | ||
| 126 | - EXPECT_EQ(node_info.broadcast_node_params.src_shape[0].expr, ge::Symbol(1)); | ||
| 127 | - EXPECT_EQ(node_info.broadcast_node_params.dst_shape[1].role, ascir_param::ParamExprRole::kActualSize); | ||
| 128 | -} | ||
| 129 | ascir::FusedScheduledResult BuildGatherReduceScheduleResult(const af::AscGraph &gather_graph, | 96 | ascir::FusedScheduledResult BuildGatherReduceScheduleResult(const af::AscGraph &gather_graph, |
| 130 | const af::AscGraph &reduce_graph, | 97 | const af::AscGraph &reduce_graph, |
| 131 | const bool enable_group_parallel) { | 98 | const bool enable_group_parallel) { |
| @@ -101,96 +101,6 @@ TEST(AscirNodeParamsTest, EnrichGraphRegistersVectorFuncParams) { | |||
| 101 | EXPECT_TRUE(stored->output_dims.empty()); | 101 | EXPECT_TRUE(stored->output_dims.empty()); |
| 102 | } | 102 | } |
| 103 | 103 | ||
| 104 | -TEST(AscirNodeParamsTest, BroadcastParamsCanBeStoredInAscirNodeParams) { | ||
| 105 | - ascir_param::AscirNodeParams params; | ||
| 106 | - params.specific_params = ascir_param::BroadcastNodeParams{}; | ||
| 107 | - | ||
| 108 | - const auto *stored = ascir_param::GetSpecificParams<ascir_param::BroadcastNodeParams>(params); | ||
| 109 | - ASSERT_NE(stored, nullptr); | ||
| 110 | - EXPECT_FALSE(stored->valid); | ||
| 111 | -} | ||
| 112 | - | ||
| 113 | -TEST(AscirNodeParamsTest, EnrichGraphRegistersBroadcastParams) { | ||
| 114 | - af::AscGraph graph("test_graph"); | ||
| 115 | - af::ascir_op::Broadcast broadcast_op("broadcast"); | ||
| 116 | - graph.AddNode(broadcast_op); | ||
| 117 | - auto node = graph.FindNode("broadcast"); | ||
| 118 | - ASSERT_NE(node, nullptr); | ||
| 119 | - | ||
| 120 | - ExpectEnrichSuccess(graph); | ||
| 121 | - auto params = ascir_param::GetAscirNodeParams(node); | ||
| 122 | - ASSERT_NE(params, nullptr); | ||
| 123 | - EXPECT_EQ(params->api_name, "Broadcast"); | ||
| 124 | - EXPECT_EQ(params->status, ascir_param::ParamBuildStatus::kBuilt); | ||
| 125 | - const auto *stored = ascir_param::GetSpecificParams<ascir_param::BroadcastNodeParams>(*params); | ||
| 126 | - ASSERT_NE(stored, nullptr); | ||
| 127 | - EXPECT_FALSE(stored->valid); | ||
| 128 | -} | ||
| 129 | - | ||
| 130 | -TEST(AscirNodeParamsTest, EnrichGraphPreservesBroadcastParamsAcrossRepeatedCalls) { | ||
| 131 | - af::AscGraph graph("test_graph"); | ||
| 132 | - af::ascir_op::Broadcast broadcast_op("broadcast"); | ||
| 133 | - graph.AddNode(broadcast_op); | ||
| 134 | - auto node = graph.FindNode("broadcast"); | ||
| 135 | - ASSERT_NE(node, nullptr); | ||
| 136 | - | ||
| 137 | - auto params = std::make_shared<ascir_param::AscirNodeParams>(); | ||
| 138 | - params->api_name = "Broadcast"; | ||
| 139 | - params->status = ascir_param::ParamBuildStatus::kBuilt; | ||
| 140 | - ascir_param::BroadcastNodeParams broadcast; | ||
| 141 | - broadcast.valid = true; | ||
| 142 | - broadcast.is_scalar = true; | ||
| 143 | - broadcast.duplicate_count = Expr(8); | ||
| 144 | - broadcast.src_shape = {{Expr(1), ascir_param::ParamExprRole::kSemantic}, | ||
| 145 | - {Expr(8), ascir_param::ParamExprRole::kActualSize}}; | ||
| 146 | - broadcast.dst_shape = {{Expr(4), ascir_param::ParamExprRole::kActualSize}, | ||
| 147 | - {Expr(8), ascir_param::ParamExprRole::kActualSize}}; | ||
| 148 | - params->specific_params = broadcast; | ||
| 149 | - ASSERT_TRUE(node->GetOpDesc()->SetExtAttr("AscirNodeParams", params)); | ||
| 150 | - | ||
| 151 | - ExpectEnrichSuccess(graph); | ||
| 152 | - ExpectEnrichSuccess(graph); | ||
| 153 | - params = ascir_param::GetAscirNodeParams(node); | ||
| 154 | - ASSERT_NE(params, nullptr); | ||
| 155 | - const auto *stored = ascir_param::GetSpecificParams<ascir_param::BroadcastNodeParams>(*params); | ||
| 156 | - ASSERT_NE(stored, nullptr); | ||
| 157 | - EXPECT_TRUE(stored->valid); | ||
| 158 | - EXPECT_TRUE(stored->is_scalar); | ||
| 159 | - EXPECT_EQ(stored->duplicate_count, Expr(8)); | ||
| 160 | - EXPECT_EQ(stored->src_shape[0].expr, Expr(1)); | ||
| 161 | - EXPECT_EQ(stored->dst_shape[1].role, ascir_param::ParamExprRole::kActualSize); | ||
| 162 | -} | ||
| 163 | - | ||
| 164 | -TEST(AscirNodeParamsTest, EnrichGraphPreservesPrebuiltBroadcastParams) { | ||
| 165 | - af::AscGraph graph("test_graph"); | ||
| 166 | - af::ascir_op::Broadcast broadcast_op("broadcast"); | ||
| 167 | - graph.AddNode(broadcast_op); | ||
| 168 | - auto node = graph.FindNode("broadcast"); | ||
| 169 | - ASSERT_NE(node, nullptr); | ||
| 170 | - | ||
| 171 | - auto params = std::make_shared<ascir_param::AscirNodeParams>(); | ||
| 172 | - params->api_name = "Broadcast"; | ||
| 173 | - params->status = ascir_param::ParamBuildStatus::kBuilt; | ||
| 174 | - ascir_param::BroadcastNodeParams broadcast; | ||
| 175 | - broadcast.valid = true; | ||
| 176 | - broadcast.duplicate_count = Expr(16); | ||
| 177 | - broadcast.src_shape = {{Expr(1), ascir_param::ParamExprRole::kSemantic}, | ||
| 178 | - {Expr(16), ascir_param::ParamExprRole::kSize}}; | ||
| 179 | - broadcast.dst_shape = {{Expr(2), ascir_param::ParamExprRole::kSize}, {Expr(16), ascir_param::ParamExprRole::kSize}}; | ||
| 180 | - params->specific_params = broadcast; | ||
| 181 | - ASSERT_TRUE(node->GetOpDesc()->SetExtAttr("AscirNodeParams", params)); | ||
| 182 | - | ||
| 183 | - ExpectEnrichSuccess(graph); | ||
| 184 | - params = ascir_param::GetAscirNodeParams(node); | ||
| 185 | - ASSERT_NE(params, nullptr); | ||
| 186 | - const auto *stored = ascir_param::GetSpecificParams<ascir_param::BroadcastNodeParams>(*params); | ||
| 187 | - ASSERT_NE(stored, nullptr); | ||
| 188 | - EXPECT_TRUE(stored->valid); | ||
| 189 | - EXPECT_EQ(stored->duplicate_count, Expr(16)); | ||
| 190 | - EXPECT_EQ(stored->src_shape[1].expr, Expr(16)); | ||
| 191 | - EXPECT_EQ(stored->dst_shape[0].role, ascir_param::ParamExprRole::kSize); | ||
| 192 | -} | ||
| 193 | - | ||
| 194 | TEST(AscirNodeParamsTest, EnrichReduceParamsForArSingleReduce) { | 104 | TEST(AscirNodeParamsTest, EnrichReduceParamsForArSingleReduce) { |
| 195 | auto env = MakeArReduceEnv("max"); | 105 | auto env = MakeArReduceEnv("max"); |
| 196 | 106 | ||
| @@ -63,8 +63,6 @@ TEST_F(TestBackendScalarBrcE2e, ScalarBrcE2eCodegen) { | |||
| 63 | EXPECT_EQ(optimizer.Optimize(graph, fused_schedule_result), 0); | 63 | EXPECT_EQ(optimizer.Optimize(graph, fused_schedule_result), 0); |
| 64 | codegen::CodegenResult result; | 64 | codegen::CodegenResult result; |
| 65 | EXPECT_EQ(codegen.Generate(shape_info, fused_schedule_result, result), 0); | 65 | EXPECT_EQ(codegen.Generate(shape_info, fused_schedule_result, result), 0); |
| 66 | - EXPECT_EQ(result.tiling.find("local_3_actual_size"), std::string::npos); | ||
| 67 | - EXPECT_NE(result.tiling.find("z0z1t_size"), std::string::npos); | ||
| 68 | kernel_file << tilig_stub << RemoveSubDirInclude(result.kernel); | 66 | kernel_file << tilig_stub << RemoveSubDirInclude(result.kernel); |
| 69 | tiling_file << result.tiling; | 67 | tiling_file << result.tiling; |
| 70 | tiling_data_file << result.tiling_data; | 68 | tiling_data_file << result.tiling_data; |
Dautofuse/tests/v35/ut/att/gen_model_info/api_perf_register/test_broadcast_last_axis_perf_v2.cpp+0-464
| @@ -1,464 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 of the License. | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | -namespace att { | ||
| 18 | -namespace { | ||
| 19 | -using ascendcapi_v2::LastAxisBranch; | ||
| 20 | -using ascendcapi_v2::VfCostAccumulator; | ||
| 21 | - | ||
| 22 | -NodeDetail MakeLastNode(const std::string &dtype, const std::vector<int64_t> &src, const std::vector<int64_t> &dst) { | ||
| 23 | - NodeDetail node; | ||
| 24 | - node.broadcast_node_params.valid = true; | ||
| 25 | - node.broadcast_node_params.const_rank = static_cast<int32_t>(src.size()); | ||
| 26 | - for (size_t i = 0U; i < src.size(); ++i) { | ||
| 27 | - node.broadcast_node_params.src_shape.push_back({CreateExpr(src[i]), ascir_param::ParamExprRole::kSemantic}); | ||
| 28 | - node.broadcast_node_params.dst_shape.push_back({CreateExpr(dst[i]), ascir_param::ParamExprRole::kSemantic}); | ||
| 29 | - node.input_dims.push_back(CreateExpr(src[i])); | ||
| 30 | - node.output_dims.push_back(CreateExpr(dst[i])); | ||
| 31 | - node.repeats.push_back(CreateExpr(src[i])); | ||
| 32 | - } | ||
| 33 | - node.input_dtype = {dtype}; | ||
| 34 | - node.output_dtype = {dtype}; | ||
| 35 | - return node; | ||
| 36 | -} | ||
| 37 | - | ||
| 38 | -void AddExpected(const std::string &op, const std::string &dtype, int64_t count, VfCostAccumulator &acc) { | ||
| 39 | - ASSERT_EQ(ascendcapi_v2::AddVfInstructPerf(op, dtype, CreateExpr(count), 1U, acc), af::SUCCESS); | ||
| 40 | -} | ||
| 41 | - | ||
| 42 | -Expr ExpectedE2B(const std::string &dtype, int64_t count) { | ||
| 43 | - VfCostAccumulator acc; | ||
| 44 | - AddExpected(kUpdateMask, dtype, count, acc); | ||
| 45 | - AddExpected(kLoad, dtype, count, acc); | ||
| 46 | - AddExpected(kStore, dtype, count, acc); | ||
| 47 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 48 | -} | ||
| 49 | - | ||
| 50 | -Expr ExpectedAligned(const std::string &dtype, int64_t loads, int64_t updates, int64_t stores) { | ||
| 51 | - VfCostAccumulator acc; | ||
| 52 | - AddExpected(kUpdateMask, dtype, updates, acc); | ||
| 53 | - AddExpected(kLoad, dtype, loads, acc); | ||
| 54 | - AddExpected(kStore, dtype, stores, acc); | ||
| 55 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 56 | -} | ||
| 57 | - | ||
| 58 | -Expr ExpectedRankTwoLargerAligned() { | ||
| 59 | - return ExpectedAligned(kFloat16, 2, 4, 4); | ||
| 60 | -} | ||
| 61 | - | ||
| 62 | -Expr ExpectedRankThreeLargerAligned() { | ||
| 63 | - return ExpectedAligned(kFloat16, 6, 12, 12); | ||
| 64 | -} | ||
| 65 | - | ||
| 66 | -Expr ExpectedRankFourLargerAligned() { | ||
| 67 | - return ExpectedAligned(kFloat16, 24, 8, 24); | ||
| 68 | -} | ||
| 69 | - | ||
| 70 | -Expr ExpectedUnaligned(const std::string &dtype, int64_t loads, int64_t stores) { | ||
| 71 | - VfCostAccumulator acc; | ||
| 72 | - AddExpected(kLoad, dtype, loads, acc); | ||
| 73 | - AddExpected(kStore, dtype, stores, acc); | ||
| 74 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, dtype, acc.max_latency, acc.throughput, CreateExpr(1)), | ||
| 75 | - af::SUCCESS); | ||
| 76 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 77 | -} | ||
| 78 | - | ||
| 79 | -Expr ExpectedDynamicLast(const std::string &dtype, int64_t outer, int64_t stores, int64_t helper_calls) { | ||
| 80 | - VfCostAccumulator acc; | ||
| 81 | - AddExpected(kLoad, dtype, outer, acc); | ||
| 82 | - AddExpected(kStore, dtype, stores, acc); | ||
| 83 | - std::string helper_dtype; | ||
| 84 | - EXPECT_EQ(ascendcapi_v2::GetEffectiveHelperDtype(dtype, helper_dtype), af::SUCCESS); | ||
| 85 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, helper_dtype, acc.max_latency, acc.throughput, | ||
| 86 | - CreateExpr(helper_calls)), | ||
| 87 | - af::SUCCESS); | ||
| 88 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 89 | -} | ||
| 90 | - | ||
| 91 | -Expr ExpectedRankTwoGather(bool two) { | ||
| 92 | - VfCostAccumulator acc; | ||
| 93 | - AddExpected(kDuplicate, kInt16, two ? 2 : 1, acc); | ||
| 94 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kInt16, acc.max_latency, acc.throughput, CreateExpr(1)), | ||
| 95 | - af::SUCCESS); | ||
| 96 | - AddExpected(kDiv, kInt16, 1, acc); | ||
| 97 | - AddExpected(kStore, kInt16, 1, acc); | ||
| 98 | - if (two) { | ||
| 99 | - AddExpected(kMuls, kUInt16, 1, acc); | ||
| 100 | - AddExpected(kAdd, kUInt16, 1, acc); | ||
| 101 | - AddExpected(kAdds, kUInt16, 1, acc); | ||
| 102 | - } else { | ||
| 103 | - AddExpected(kUpdateMask, kFloat16, 1, acc); | ||
| 104 | - } | ||
| 105 | - AddExpected(kLoad, kUInt16, 1, acc); | ||
| 106 | - EXPECT_EQ( | ||
| 107 | - VfPerfUtils::AddVfInstructPerf(kPlaceholder, kUInt16, acc.max_latency, acc.throughput, CreateExpr(two ? 2 : 1)), | ||
| 108 | - af::SUCCESS); | ||
| 109 | - AddExpected(kStore, kFloat16, two ? 2 : 1, acc); | ||
| 110 | - if (two) { | ||
| 111 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kFloat16, acc.max_latency, acc.throughput, CreateExpr(1)), | ||
| 112 | - af::SUCCESS); | ||
| 113 | - } | ||
| 114 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 115 | -} | ||
| 116 | - | ||
| 117 | -Expr ExpectedGatherWrapper(size_t rank, int64_t arithmetic, int64_t gathers) { | ||
| 118 | - VfCostAccumulator acc; | ||
| 119 | - AddExpected(kDuplicate, kInt16, static_cast<int64_t>(rank * 2U), acc); | ||
| 120 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kInt16, acc.max_latency, acc.throughput, CreateExpr(1)), | ||
| 121 | - af::SUCCESS); | ||
| 122 | - AddExpected(kDiv, kInt16, static_cast<int64_t>(rank), acc); | ||
| 123 | - AddExpected(kMul, kInt16, static_cast<int64_t>(rank + 1U), acc); | ||
| 124 | - AddExpected(kSub, kInt16, static_cast<int64_t>(rank), acc); | ||
| 125 | - AddExpected(kMulAddDst, kInt16, static_cast<int64_t>(rank - 1U), acc); | ||
| 126 | - AddExpected(kStore, kInt16, 1, acc); | ||
| 127 | - AddExpected(kDuplicate, kInt16, 2, acc); | ||
| 128 | - AddExpected(kLoad, kInt16, 1, acc); | ||
| 129 | - AddExpected(kMuls, kUInt16, arithmetic, acc); | ||
| 130 | - AddExpected(kAdd, kUInt16, arithmetic, acc); | ||
| 131 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kUInt16, acc.max_latency, acc.throughput, CreateExpr(gathers)), | ||
| 132 | - af::SUCCESS); | ||
| 133 | - AddExpected(kStore, kFloat16, gathers, acc); | ||
| 134 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kFloat16, acc.max_latency, acc.throughput, CreateExpr(1)), | ||
| 135 | - af::SUCCESS); | ||
| 136 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 137 | -} | ||
| 138 | - | ||
| 139 | -Expr ExpectedRankTwoB8GatherOne() { | ||
| 140 | - VfCostAccumulator acc; | ||
| 141 | - AddExpected(kDuplicate, kInt16, 2, acc); | ||
| 142 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kInt16, acc.max_latency, acc.throughput, CreateExpr(2)), | ||
| 143 | - af::SUCCESS); | ||
| 144 | - AddExpected(kDiv, kInt16, 2, acc); | ||
| 145 | - AddExpected(kStore, kInt16, 2, acc); | ||
| 146 | - AddExpected(kUpdateMask, kUInt8, 1, acc); | ||
| 147 | - AddExpected(kLoad, kUInt16, 2, acc); | ||
| 148 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kUInt16, acc.max_latency, acc.throughput, CreateExpr(2)), | ||
| 149 | - af::SUCCESS); | ||
| 150 | - AddExpected(kDeInterleave, kUInt8, 1, acc); | ||
| 151 | - AddExpected(kStore, kUInt8, 1, acc); | ||
| 152 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 153 | -} | ||
| 154 | - | ||
| 155 | -Expr ExpectedRankThreeB8GatherWrapper() { | ||
| 156 | - VfCostAccumulator acc; | ||
| 157 | - AddExpected(kDuplicate, kInt16, 6, acc); | ||
| 158 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kInt16, acc.max_latency, acc.throughput, CreateExpr(2)), | ||
| 159 | - af::SUCCESS); | ||
| 160 | - AddExpected(kDiv, kInt16, 6, acc); | ||
| 161 | - AddExpected(kMul, kInt16, 8, acc); | ||
| 162 | - AddExpected(kSub, kInt16, 6, acc); | ||
| 163 | - AddExpected(kMulAddDst, kInt16, 4, acc); | ||
| 164 | - AddExpected(kStore, kInt16, 2, acc); | ||
| 165 | - AddExpected(kDuplicate, kInt16, 2, acc); | ||
| 166 | - AddExpected(kLoad, kInt16, 2, acc); | ||
| 167 | - AddExpected(kMuls, kUInt16, 2, acc); | ||
| 168 | - AddExpected(kAdd, kUInt16, 4, acc); | ||
| 169 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kUInt16, acc.max_latency, acc.throughput, CreateExpr(2)), | ||
| 170 | - af::SUCCESS); | ||
| 171 | - AddExpected(kDeInterleave, kUInt8, 1, acc); | ||
| 172 | - AddExpected(kStore, kUInt8, 1, acc); | ||
| 173 | - EXPECT_EQ(VfPerfUtils::AddVfInstructPerf(kPlaceholder, kUInt8, acc.max_latency, acc.throughput, CreateExpr(1)), | ||
| 174 | - af::SUCCESS); | ||
| 175 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 176 | -} | ||
| 177 | - | ||
| 178 | -ascendcapi_v2::BroadcastTilingInfo BuildTiling(const NodeDetail &node) { | ||
| 179 | - ascendcapi_v2::ParamExprInputs src_inputs{node.input_dims, node.input_dims, node.repeats}; | ||
| 180 | - ascendcapi_v2::ParamExprInputs dst_inputs{node.output_dims, node.output_dims, node.output_dims}; | ||
| 181 | - ascendcapi_v2::BroadcastTilingInfo tiling; | ||
| 182 | - EXPECT_EQ( | ||
| 183 | - ascendcapi_v2::BuildBroadcastTiling(node.broadcast_node_params.src_shape, node.broadcast_node_params.dst_shape, | ||
| 184 | - node.input_dtype[0], src_inputs, dst_inputs, tiling), | ||
| 185 | - af::SUCCESS); | ||
| 186 | - return tiling; | ||
| 187 | -} | ||
| 188 | - | ||
| 189 | -Expr ActualCost(const NodeDetail &node) { | ||
| 190 | - PerfOutputInfo perf; | ||
| 191 | - EXPECT_EQ(ascendcperf_v2::BroadcastPerf(node, perf), af::SUCCESS); | ||
| 192 | - return perf.pipe_res[PipeType::AIV_VEC]; | ||
| 193 | -} | ||
| 194 | - | ||
| 195 | -Expr ActualLastCost(const NodeDetail &node, const ascendcapi_v2::BroadcastTilingInfo &tiling) { | ||
| 196 | - PerfOutputInfo perf; | ||
| 197 | - EXPECT_EQ(ascendcperf_v2::BuildLastAxisPerf(node, tiling, perf), af::SUCCESS); | ||
| 198 | - return perf.pipe_res[PipeType::AIV_VEC]; | ||
| 199 | -} | ||
| 200 | - | ||
| 201 | -struct LastBranchCase { | ||
| 202 | - std::string dtype; | ||
| 203 | - std::vector<int64_t> src; | ||
| 204 | - std::vector<int64_t> dst; | ||
| 205 | - LastAxisBranch branch; | ||
| 206 | -}; | ||
| 207 | - | ||
| 208 | -struct FixedVectorSpec { | ||
| 209 | - int64_t vl; | ||
| 210 | - int64_t half_vl; | ||
| 211 | - int64_t block; | ||
| 212 | -}; | ||
| 213 | - | ||
| 214 | -FixedVectorSpec GetFixedVectorSpec(const std::string &dtype) { | ||
| 215 | - if (dtype == kUInt8 || dtype == kInt8) { | ||
| 216 | - return {256, 128, 32}; | ||
| 217 | - } | ||
| 218 | - if (dtype == kFloat32 || dtype == kUInt32 || dtype == kInt32) { | ||
| 219 | - return {64, 32, 8}; | ||
| 220 | - } | ||
| 221 | - return {128, 64, 16}; | ||
| 222 | -} | ||
| 223 | - | ||
| 224 | -int64_t Product(const std::vector<int64_t> &shape) { | ||
| 225 | - int64_t result = 1; | ||
| 226 | - for (const auto dim : shape) { | ||
| 227 | - result *= dim; | ||
| 228 | - } | ||
| 229 | - return result; | ||
| 230 | -} | ||
| 231 | - | ||
| 232 | -// Independent loop specification copied from the source helper conditions; it is not a production classifier oracle. | ||
| 233 | -LastAxisBranch FixedLastAxisBranch(const std::string &dtype, const std::vector<int64_t> &src, | ||
| 234 | - const std::vector<int64_t> &dst, int32_t const_rank) { | ||
| 235 | - const auto spec = GetFixedVectorSpec(dtype); | ||
| 236 | - if (const_rank > 4) { | ||
| 237 | - return dst.back() <= spec.vl ? LastAxisBranch::kDynamicLessThanVlUnaligned | ||
| 238 | - : LastAxisBranch::kDynamicLargerThanVlUnaligned; | ||
| 239 | - } | ||
| 240 | - size_t rank = dst.size(); | ||
| 241 | - size_t first_dim = 0U; | ||
| 242 | - while (rank > 1U && dst[first_dim] == 1) { | ||
| 243 | - ++first_dim; | ||
| 244 | - --rank; | ||
| 245 | - } | ||
| 246 | - const int64_t last = dst.back(); | ||
| 247 | - const bool b8 = dtype == kUInt8 || dtype == kInt8; | ||
| 248 | - const bool aligned = last % spec.block == 0; | ||
| 249 | - const bool second_axis_contiguous = src.size() > 1U && src[src.size() - 2U] != 1; | ||
| 250 | - const int64_t dst_size = Product(dst); | ||
| 251 | - if (rank == 2U && last == spec.block && !b8) { | ||
| 252 | - return LastAxisBranch::kE2B; | ||
| 253 | - } | ||
| 254 | - if (rank == 3U && !b8 && second_axis_contiguous && last == spec.block && dst[dst.size() - 2U] * last > spec.half_vl && | ||
| 255 | - dst[dst.size() - 2U] % (spec.vl / spec.block) == 0) { | ||
| 256 | - return dst[dst.size() - 2U] * last > spec.vl ? LastAxisBranch::kE2BLargerThanVl : LastAxisBranch::kE2BLessThanVl; | ||
| 257 | - } | ||
| 258 | - if (rank == 4U && !b8 && src.size() > 2U && src[src.size() - 2U] != 1 && last == spec.block && | ||
| 259 | - dst[dst.size() - 2U] % (spec.vl / spec.block) == 0) { | ||
| 260 | - return LastAxisBranch::kE2B; | ||
| 261 | - } | ||
| 262 | - if (last < spec.half_vl && (rank == 2U || !b8)) { | ||
| 263 | - if (rank == 2U) { | ||
| 264 | - return dst_size < spec.vl ? LastAxisBranch::kGatherOne : LastAxisBranch::kGatherTwo; | ||
| 265 | - } | ||
| 266 | - return rank == 3U ? LastAxisBranch::kGatherWrapper : LastAxisBranch::kGatherWrapperForFourDim; | ||
| 267 | - } | ||
| 268 | - if (last <= spec.vl) { | ||
| 269 | - if (rank == 2U) { | ||
| 270 | - return LastAxisBranch::kLessThanVlUnaligned; | ||
| 271 | - } | ||
| 272 | - if (rank == 3U) { | ||
| 273 | - return aligned ? LastAxisBranch::kLessThanVlAligned : LastAxisBranch::kLessThanVlUnaligned; | ||
| 274 | - } | ||
| 275 | - if (aligned) { | ||
| 276 | - return LastAxisBranch::kLessThanVlAligned; | ||
| 277 | - } | ||
| 278 | - return const_rank == -1 ? LastAxisBranch::kFallback : LastAxisBranch::kLessThanVlUnaligned; | ||
| 279 | - } | ||
| 280 | - if (aligned) { | ||
| 281 | - return LastAxisBranch::kLargerThanVlAligned; | ||
| 282 | - } | ||
| 283 | - return const_rank == -1 ? LastAxisBranch::kFallback : LastAxisBranch::kLargerThanVlUnaligned; | ||
| 284 | -} | ||
| 285 | - | ||
| 286 | -TEST(BroadcastLastAxisPerfV2, RankTwoCoversImplementationLeavesAndBoundaries) { | ||
| 287 | - const std::vector<LastBranchCase> cases = { | ||
| 288 | - {kFloat16, {2, 1}, {2, 16}, LastAxisBranch::kE2B}, | ||
| 289 | - {kFloat16, {2, 1}, {2, 63}, LastAxisBranch::kGatherOne}, | ||
| 290 | - {kFloat16, {15, 1}, {15, 8}, LastAxisBranch::kGatherOne}, | ||
| 291 | - {kFloat16, {16, 1}, {16, 8}, LastAxisBranch::kGatherTwo}, | ||
| 292 | - {kFloat16, {17, 1}, {17, 8}, LastAxisBranch::kGatherTwo}, | ||
| 293 | - {kFloat16, {8, 1}, {8, 16}, LastAxisBranch::kE2B}, | ||
| 294 | - {kFloat16, {9, 1}, {9, 15}, LastAxisBranch::kGatherTwo}, | ||
| 295 | - {kFloat16, {2, 1}, {2, 17}, LastAxisBranch::kGatherOne}, | ||
| 296 | - {kFloat16, {2, 1}, {2, 64}, LastAxisBranch::kLessThanVlUnaligned}, | ||
| 297 | - {kFloat16, {2, 1}, {2, 65}, LastAxisBranch::kLessThanVlUnaligned}, | ||
| 298 | - {kFloat16, {2, 1}, {2, 127}, LastAxisBranch::kLessThanVlUnaligned}, | ||
| 299 | - {kFloat16, {2, 1}, {2, 128}, LastAxisBranch::kLessThanVlUnaligned}, | ||
| 300 | - {kFloat16, {2, 1}, {2, 144}, LastAxisBranch::kLargerThanVlAligned}, | ||
| 301 | - {kFloat16, {2, 1}, {2, 129}, LastAxisBranch::kLargerThanVlUnaligned}, | ||
| 302 | - }; | ||
| 303 | - for (const auto &item : cases) { | ||
| 304 | - const auto node = MakeLastNode(item.dtype, item.src, item.dst); | ||
| 305 | - EXPECT_EQ(FixedLastAxisBranch(item.dtype, item.src, item.dst, static_cast<int32_t>(item.dst.size())), item.branch); | ||
| 306 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(node), item.branch); | ||
| 307 | - } | ||
| 308 | -} | ||
| 309 | - | ||
| 310 | -TEST(BroadcastLastAxisPerfV2, RankThreeCoversBothE2BPathsAndGatherWrapper) { | ||
| 311 | - const std::vector<LastBranchCase> cases = { | ||
| 312 | - {kFloat16, {2, 8, 1}, {2, 8, 16}, LastAxisBranch::kE2BLessThanVl}, | ||
| 313 | - {kFloat16, {2, 16, 1}, {2, 16, 16}, LastAxisBranch::kE2BLargerThanVl}, | ||
| 314 | - {kFloat16, {2, 3, 1}, {2, 3, 15}, LastAxisBranch::kGatherWrapper}, | ||
| 315 | - {kFloat16, {2, 3, 1}, {2, 3, 64}, LastAxisBranch::kLessThanVlAligned}, | ||
| 316 | - {kFloat16, {2, 3, 1}, {2, 3, 65}, LastAxisBranch::kLessThanVlUnaligned}, | ||
| 317 | - {kFloat16, {2, 3, 1}, {2, 3, 144}, LastAxisBranch::kLargerThanVlAligned}, | ||
| 318 | - {kFloat16, {2, 3, 1}, {2, 3, 129}, LastAxisBranch::kLargerThanVlUnaligned}, | ||
| 319 | - }; | ||
| 320 | - for (const auto &item : cases) { | ||
| 321 | - const auto node = MakeLastNode(item.dtype, item.src, item.dst); | ||
| 322 | - EXPECT_EQ(FixedLastAxisBranch(item.dtype, item.src, item.dst, static_cast<int32_t>(item.dst.size())), item.branch); | ||
| 323 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(node), item.branch); | ||
| 324 | - } | ||
| 325 | -} | ||
| 326 | - | ||
| 327 | -TEST(BroadcastLastAxisPerfV2, RankThreeDynamicRankKeepsAlignedStaticAndFallsBackOnUnaligned) { | ||
| 328 | - auto unaligned_node = MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 65}); | ||
| 329 | - unaligned_node.broadcast_node_params.const_rank = -1; | ||
| 330 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(unaligned_node), LastAxisBranch::kDynamicLessThanVlUnaligned); | ||
| 331 | - | ||
| 332 | - auto aligned_node = MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 64}); | ||
| 333 | - aligned_node.broadcast_node_params.const_rank = -1; | ||
| 334 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(aligned_node), LastAxisBranch::kLessThanVlAligned); | ||
| 335 | -} | ||
| 336 | - | ||
| 337 | -TEST(BroadcastLastAxisPerfV2, RankFourFoldedToThreeDimAlignedUsesAlignedBranch) { | ||
| 338 | - const auto node = MakeLastNode(kFloat16, {1, 1, 2, 1}, {1, 1, 2, 64}); | ||
| 339 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(node), LastAxisBranch::kLessThanVlAligned); | ||
| 340 | -} | ||
| 341 | - | ||
| 342 | -TEST(BroadcastLastAxisPerfV2, RankFourCoversWrapperLeavesAndReductionToThreeDim) { | ||
| 343 | - const std::vector<LastBranchCase> cases = { | ||
| 344 | - {kFloat16, {2, 3, 8, 1}, {2, 3, 8, 16}, LastAxisBranch::kE2B}, | ||
| 345 | - {kFloat16, {2, 3, 2, 1}, {2, 3, 2, 15}, LastAxisBranch::kGatherWrapperForFourDim}, | ||
| 346 | - {kFloat16, {2, 3, 2, 1}, {2, 3, 2, 64}, LastAxisBranch::kLessThanVlAligned}, | ||
| 347 | - {kFloat16, {2, 3, 2, 1}, {2, 3, 2, 65}, LastAxisBranch::kLessThanVlUnaligned}, | ||
| 348 | - {kFloat16, {2, 3, 2, 1}, {2, 3, 2, 144}, LastAxisBranch::kLargerThanVlAligned}, | ||
| 349 | - {kFloat16, {2, 3, 2, 1}, {2, 3, 2, 129}, LastAxisBranch::kLargerThanVlUnaligned}, | ||
| 350 | - {kFloat16, {1, 2, 3, 1}, {1, 2, 3, 65}, LastAxisBranch::kLessThanVlUnaligned}, | ||
| 351 | - }; | ||
| 352 | - for (const auto &item : cases) { | ||
| 353 | - const auto node = MakeLastNode(item.dtype, item.src, item.dst); | ||
| 354 | - EXPECT_EQ(FixedLastAxisBranch(item.dtype, item.src, item.dst, static_cast<int32_t>(item.dst.size())), item.branch); | ||
| 355 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(node), item.branch); | ||
| 356 | - } | ||
| 357 | -} | ||
| 358 | - | ||
| 359 | -TEST(BroadcastLastAxisPerfV2, LeafHelpersUseExactE2BAndAlignedCounts) { | ||
| 360 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 8, 1}, {2, 8, 16})), ExpectedE2B(kFloat16, 2)); | ||
| 361 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 16, 1}, {2, 16, 16})), ExpectedE2B(kFloat16, 4)); | ||
| 362 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 8, 1}, {2, 3, 8, 16})), ExpectedE2B(kFloat16, 6)); | ||
| 363 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 2, 1}, {2, 3, 2, 144})), ExpectedRankFourLargerAligned()); | ||
| 364 | -} | ||
| 365 | - | ||
| 366 | -TEST(BroadcastLastAxisPerfV2, LeafHelpersUseExactUnalignedCounts) { | ||
| 367 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 1}, {2, 65})), ExpectedUnaligned(kFloat16, 2, 2)); | ||
| 368 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 65})), ExpectedUnaligned(kFloat16, 6, 6)); | ||
| 369 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 2, 1}, {2, 3, 2, 65})), ExpectedUnaligned(kFloat16, 12, 12)); | ||
| 370 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {1, 2, 3, 1}, {1, 2, 3, 65})), ExpectedUnaligned(kFloat16, 6, 6)); | ||
| 371 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 129})), ExpectedUnaligned(kFloat16, 6, 12)); | ||
| 372 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 2, 1}, {2, 3, 2, 129})), ExpectedUnaligned(kFloat16, 24, 24)); | ||
| 373 | -} | ||
| 374 | - | ||
| 375 | -TEST(BroadcastLastAxisPerfV2, RankTwoGatherAndAlignedHelpersUseExactCounts) { | ||
| 376 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 1}, {2, 15})), ExpectedRankTwoGather(false)); | ||
| 377 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {9, 1}, {9, 15})), ExpectedRankTwoGather(true)); | ||
| 378 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 1}, {2, 144})), ExpectedRankTwoLargerAligned()); | ||
| 379 | - EXPECT_EQ(ActualCost(MakeLastNode(kUInt8, {2, 1}, {2, 63})), ExpectedRankTwoB8GatherOne()); | ||
| 380 | -} | ||
| 381 | - | ||
| 382 | -TEST(BroadcastLastAxisPerfV2, RankThreeAndFourHelpersUseExactCounts) { | ||
| 383 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 15})), ExpectedGatherWrapper(3U, 3, 2)); | ||
| 384 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 2, 1}, {2, 3, 2, 15})), ExpectedGatherWrapper(4U, 5, 3)); | ||
| 385 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 64})), ExpectedAligned(kFloat16, 6, 1, 6)); | ||
| 386 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 144})), ExpectedRankThreeLargerAligned()); | ||
| 387 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 2, 1}, {2, 3, 2, 64})), ExpectedAligned(kFloat16, 12, 1, 12)); | ||
| 388 | - EXPECT_EQ(ActualCost(MakeLastNode(kUInt8, {2, 3, 1}, {2, 3, 15})), ExpectedRankThreeB8GatherWrapper()); | ||
| 389 | -} | ||
| 390 | - | ||
| 391 | -TEST(BroadcastLastAxisPerfV2, RankTwoThreeFourLargerAlignedUseDifferentLoadCounts) { | ||
| 392 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 1}, {2, 144})), ExpectedRankTwoLargerAligned()); | ||
| 393 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 1}, {2, 3, 144})), ExpectedRankThreeLargerAligned()); | ||
| 394 | - EXPECT_EQ(ActualCost(MakeLastNode(kFloat16, {2, 3, 2, 1}, {2, 3, 2, 144})), ExpectedRankFourLargerAligned()); | ||
| 395 | -} | ||
| 396 | - | ||
| 397 | -TEST(BroadcastLastAxisPerfV2, DynamicRanksFiveThroughNineUseOriginalPrefixForHelperCalls) { | ||
| 398 | - const std::vector<std::vector<int64_t>> dst_shapes = { | ||
| 399 | - {2, 3, 2, 3, 65}, {2, 3, 2, 3, 2, 65}, {2, 3, 2, 3, 2, 3, 65}, | ||
| 400 | - {2, 3, 2, 3, 2, 3, 2, 65}, {2, 3, 2, 3, 2, 3, 2, 3, 65}, | ||
| 401 | - }; | ||
| 402 | - for (const auto &dst : dst_shapes) { | ||
| 403 | - auto src = dst; | ||
| 404 | - for (size_t i = 0U; i < src.size(); ++i) { | ||
| 405 | - if (i % 2U == (src.size() - 1U) % 2U) { | ||
| 406 | - src[i] = 1; | ||
| 407 | - } | ||
| 408 | - } | ||
| 409 | - const auto node = MakeLastNode(kFloat16, src, dst); | ||
| 410 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(node), LastAxisBranch::kDynamicLessThanVlUnaligned); | ||
| 411 | - int64_t outer = 1; | ||
| 412 | - for (size_t i = 0U; i + 1U < dst.size(); ++i) { | ||
| 413 | - outer *= dst[i]; | ||
| 414 | - } | ||
| 415 | - const int64_t helper_calls = [&dst]() { | ||
| 416 | - int64_t result = 1; | ||
| 417 | - for (size_t i = 0U; i + 4U < dst.size(); ++i) { | ||
| 418 | - result *= dst[i]; | ||
| 419 | - } | ||
| 420 | - return result; | ||
| 421 | - }(); | ||
| 422 | - EXPECT_EQ(ActualCost(node), ExpectedDynamicLast(kFloat16, outer, outer, helper_calls)); | ||
| 423 | - } | ||
| 424 | -} | ||
| 425 | - | ||
| 426 | -TEST(BroadcastLastAxisPerfV2, DynamicOriginalRankFourUsesOneHelperCall) { | ||
| 427 | - const auto node = MakeLastNode(kUInt64, {2, 3, 1, 1}, {2, 3, 2, 4}); | ||
| 428 | - const auto tiling = BuildTiling(node); | ||
| 429 | - EXPECT_EQ(ascendcapi_v2::GetLastAxisBranch(tiling, kUInt64, 4), LastAxisBranch::kDynamicLessThanVlUnaligned); | ||
| 430 | - EXPECT_EQ(ActualLastCost(node, tiling), ExpectedDynamicLast(kUInt64, 48, 48, 1)); | ||
| 431 | -} | ||
| 432 | - | ||
| 433 | -TEST(BroadcastLastAxisPerfV2, DynamicB8CoversBothVectorSizePaths) { | ||
| 434 | - auto less = MakeLastNode(kUInt8, {1, 3, 1, 3, 1}, {2, 3, 2, 3, 256}); | ||
| 435 | - auto larger = MakeLastNode(kUInt8, {1, 3, 1, 3, 1}, {2, 3, 2, 3, 257}); | ||
| 436 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(less), LastAxisBranch::kDynamicLessThanVlUnaligned); | ||
| 437 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(larger), LastAxisBranch::kDynamicLargerThanVlUnaligned); | ||
| 438 | -} | ||
| 439 | - | ||
| 440 | -TEST(BroadcastLastAxisPerfV2, DynamicLargerHelperUsesExactLoopCounts) { | ||
| 441 | - const auto node = MakeLastNode(kFloat16, {1, 3, 1, 3, 1}, {2, 3, 2, 3, 129}); | ||
| 442 | - EXPECT_EQ(ActualCost(node), ExpectedDynamicLast(kFloat16, 36, 72, 2)); | ||
| 443 | -} | ||
| 444 | - | ||
| 445 | -TEST(BroadcastLastAxisPerfV2, B64LastAxisBecomesEffectiveB32AndAppendsDimension) { | ||
| 446 | - const auto tiling = BuildTiling(MakeLastNode(kUInt64, {2, 1}, {2, 4})); | ||
| 447 | - std::string helper_dtype; | ||
| 448 | - ASSERT_EQ(ascendcapi_v2::GetEffectiveHelperDtype(kUInt64, helper_dtype), af::SUCCESS); | ||
| 449 | - EXPECT_EQ(helper_dtype, kUInt32); | ||
| 450 | - ASSERT_EQ(tiling.rank, 3U); | ||
| 451 | - EXPECT_EQ(tiling.dst_shape[2], CreateExpr(2)); | ||
| 452 | - EXPECT_FALSE(ascendcapi_v2::IsLastAxisBroadcast(tiling)); | ||
| 453 | -} | ||
| 454 | - | ||
| 455 | -TEST(BroadcastLastAxisPerfV2, B64RankNineUsesLoopNumAndStrideNine) { | ||
| 456 | - const auto zero_stride = BuildTiling(MakeLastNode(kUInt64, {1, 3, 1, 3, 1, 3, 1, 3, 1}, {2, 3, 2, 3, 2, 3, 2, 3, 4})); | ||
| 457 | - // ApplyB64Tiling uses the original dstShape[0] as loop_num and appends its stride at index 9. | ||
| 458 | - EXPECT_EQ(zero_stride.loop_num, CreateExpr(2)); | ||
| 459 | - ASSERT_EQ(zero_stride.src_stride.size(), 10U); | ||
| 460 | - EXPECT_EQ(zero_stride.src_stride[9], CreateExpr(0)); | ||
| 461 | -} | ||
| 462 | - | ||
| 463 | -} // namespace | ||
| 464 | -} // namespace att | ||
| @@ -1,138 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 of the License. | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | -namespace att { | ||
| 19 | -namespace { | ||
| 20 | -using ascendcapi_v2::NlastAxisBranch; | ||
| 21 | - | ||
| 22 | -NodeDetail MakeSymbolicNode(const std::vector<Expr> &src, const std::vector<Expr> &dst) { | ||
| 23 | - NodeDetail node; | ||
| 24 | - node.broadcast_node_params.valid = true; | ||
| 25 | - node.broadcast_node_params.const_rank = static_cast<int32_t>(src.size()); | ||
| 26 | - for (size_t i = 0U; i < src.size(); ++i) { | ||
| 27 | - node.broadcast_node_params.src_shape.push_back({src[i], ascir_param::ParamExprRole::kSemantic}); | ||
| 28 | - node.broadcast_node_params.dst_shape.push_back({dst[i], ascir_param::ParamExprRole::kSemantic}); | ||
| 29 | - node.input_dims.push_back(src[i]); | ||
| 30 | - node.output_dims.push_back(dst[i]); | ||
| 31 | - node.repeats.push_back(src[i]); | ||
| 32 | - } | ||
| 33 | - node.input_dtype = {kFloat16}; | ||
| 34 | - node.output_dtype = {kFloat16}; | ||
| 35 | - return node; | ||
| 36 | -} | ||
| 37 | - | ||
| 38 | -TEST(BroadcastLastAxisPerfV2Symbolic, RankTwoRegistersFinalDynamicTreeVariable) { | ||
| 39 | - const auto node = MakeSymbolicNode({CreateExpr(2), CreateExpr(1)}, {CreateExpr(2), CreateExpr("last")}); | ||
| 40 | - PerfOutputInfo perf; | ||
| 41 | - ASSERT_EQ(ascendcperf_v2::BroadcastPerf(node, perf), af::SUCCESS); | ||
| 42 | - ASSERT_FALSE(perf.ternary_ops.empty()); | ||
| 43 | - const auto replacements = ConcursiveReplaceVars(perf.ternary_ops); | ||
| 44 | - EXPECT_NE(Str(perf.pipe_res.at(PipeType::AIV_VEC).Replace(replacements)).find("TernaryOp"), std::string::npos); | ||
| 45 | -} | ||
| 46 | - | ||
| 47 | -TEST(BroadcastLastAxisPerfV2Symbolic, RankThreeE2BConditionsCascadeIntoFinalVariable) { | ||
| 48 | - const auto node = MakeSymbolicNode({CreateExpr(2), CreateExpr("middle"), CreateExpr(1)}, | ||
| 49 | - {CreateExpr(2), CreateExpr("middle"), CreateExpr("last")}); | ||
| 50 | - PerfOutputInfo perf; | ||
| 51 | - ASSERT_EQ(ascendcperf_v2::BroadcastPerf(node, perf), af::SUCCESS); | ||
| 52 | - ASSERT_FALSE(perf.ternary_ops.empty()); | ||
| 53 | - const auto replacements = ConcursiveReplaceVars(perf.ternary_ops); | ||
| 54 | - const auto expression = Str(perf.pipe_res.at(PipeType::AIV_VEC).Replace(replacements)); | ||
| 55 | - EXPECT_NE(expression.find("last"), std::string::npos); | ||
| 56 | - EXPECT_NE(expression.find("TernaryOp"), std::string::npos); | ||
| 57 | -} | ||
| 58 | - | ||
| 59 | -TEST(BroadcastLastAxisPerfV2Symbolic, RankGreaterThanFourUnknownLastUsesDynamicTernaryRoute) { | ||
| 60 | - auto node = MakeSymbolicNode({CreateExpr(2), CreateExpr(1), CreateExpr(2), CreateExpr(1), CreateExpr("last")}, | ||
| 61 | - {CreateExpr(2), CreateExpr(3), CreateExpr(2), CreateExpr(3), CreateExpr("last")}); | ||
| 62 | - node.broadcast_node_params.const_rank = -1; | ||
| 63 | - const auto src_inputs = ascendcapi_v2::ParamExprInputs{node.input_dims, node.input_dims, node.repeats}; | ||
| 64 | - const auto dst_inputs = ascendcapi_v2::ParamExprInputs{node.output_dims, node.output_dims, node.output_dims}; | ||
| 65 | - ascendcapi_v2::BroadcastTilingInfo tiling; | ||
| 66 | - ASSERT_EQ( | ||
| 67 | - ascendcapi_v2::BuildBroadcastTiling(node.broadcast_node_params.src_shape, node.broadcast_node_params.dst_shape, | ||
| 68 | - kFloat16, src_inputs, dst_inputs, tiling), | ||
| 69 | - af::SUCCESS); | ||
| 70 | - EXPECT_EQ(ascendcapi_v2::GetLastAxisBranch(tiling, kFloat16, -1), ascendcapi_v2::LastAxisBranch::kFallback); | ||
| 71 | - EXPECT_EQ(ascendcperf_v2::GetLastAxisPerfBranch(node), ascendcapi_v2::LastAxisBranch::kFallback); | ||
| 72 | - PerfOutputInfo perf; | ||
| 73 | - ASSERT_EQ(ascendcperf_v2::BroadcastPerf(node, perf), af::SUCCESS); | ||
| 74 | - ASSERT_FALSE(perf.ternary_ops.empty()); | ||
| 75 | - const auto replacements = ConcursiveReplaceVars(perf.ternary_ops); | ||
| 76 | - const auto expression = Str(perf.pipe_res.at(PipeType::AIV_VEC).Replace(replacements)); | ||
| 77 | - EXPECT_NE(expression.find("last"), std::string::npos); | ||
| 78 | - EXPECT_NE(expression.find("TernaryOp"), std::string::npos); | ||
| 79 | -} | ||
| 80 | - | ||
| 81 | -TEST(BroadcastNlastAxisPerfV2Symbolic, RankTwoUnknownLastBuildsTernaryTree) { | ||
| 82 | - const auto node = MakeSymbolicNode({CreateExpr(1), CreateExpr("last")}, {CreateExpr(2), CreateExpr("last")}); | ||
| 83 | - PerfOutputInfo perf; | ||
| 84 | - ASSERT_EQ(ascendcperf_v2::BroadcastPerf(node, perf), af::SUCCESS); | ||
| 85 | - ASSERT_FALSE(perf.ternary_ops.empty()); | ||
| 86 | - const auto replacements = ConcursiveReplaceVars(perf.ternary_ops); | ||
| 87 | - const auto expression = Str(perf.pipe_res.at(PipeType::AIV_VEC).Replace(replacements)); | ||
| 88 | - EXPECT_NE(expression.find("last"), std::string::npos); | ||
| 89 | - EXPECT_NE(expression.find("TernaryOp"), std::string::npos); | ||
| 90 | -} | ||
| 91 | - | ||
| 92 | -TEST(BroadcastNlastAxisPerfV2Symbolic, RankThreeAndFourUnknownLastBuildTernaryTrees) { | ||
| 93 | - for (const size_t rank : {3U, 4U}) { | ||
| 94 | - std::vector<Expr> src(rank, CreateExpr(1)); | ||
| 95 | - std::vector<Expr> dst(rank, CreateExpr(2)); | ||
| 96 | - src.back() = CreateExpr(32); | ||
| 97 | - dst.back() = CreateExpr("last"); | ||
| 98 | - const auto node = MakeSymbolicNode(src, dst); | ||
| 99 | - PerfOutputInfo perf; | ||
| 100 | - ASSERT_EQ(ascendcperf_v2::BroadcastPerf(node, perf), af::SUCCESS); | ||
| 101 | - EXPECT_FALSE(perf.ternary_ops.empty()); | ||
| 102 | - const auto replacements = ConcursiveReplaceVars(perf.ternary_ops); | ||
| 103 | - EXPECT_NE(Str(perf.pipe_res.at(PipeType::AIV_VEC).Replace(replacements)).find("TernaryOp"), std::string::npos); | ||
| 104 | - } | ||
| 105 | -} | ||
| 106 | - | ||
| 107 | -TEST(BroadcastNlastAxisPerfV2Symbolic, RankGreaterThanFourUnknownLastBuildsDynamicTree) { | ||
| 108 | - const auto node = MakeSymbolicNode({CreateExpr(2), CreateExpr(1), CreateExpr(2), CreateExpr(1), CreateExpr("last")}, | ||
| 109 | - {CreateExpr(2), CreateExpr(3), CreateExpr(2), CreateExpr(3), CreateExpr("last")}); | ||
| 110 | - EXPECT_NE(ascendcperf_v2::GetNlastAxisPerfBranch(node), NlastAxisBranch::kFallback); | ||
| 111 | -} | ||
| 112 | - | ||
| 113 | -TEST(BroadcastNlastAxisPerfV2Symbolic, RankGreaterThanFourUnknownB64LastUsesDynamicTree) { | ||
| 114 | - auto node = MakeSymbolicNode({CreateExpr(2), CreateExpr(1), CreateExpr(2), CreateExpr(1), CreateExpr("last")}, | ||
| 115 | - {CreateExpr(2), CreateExpr(3), CreateExpr(2), CreateExpr(3), CreateExpr("last")}); | ||
| 116 | - node.input_dtype = {kUInt64}; | ||
| 117 | - node.output_dtype = {kUInt64}; | ||
| 118 | - EXPECT_NE(ascendcperf_v2::GetNlastAxisPerfBranch(node), NlastAxisBranch::kFallback); | ||
| 119 | - EXPECT_FALSE(ascendcperf_v2::IsBroadcastFallback(node)); | ||
| 120 | - PerfOutputInfo perf; | ||
| 121 | - ASSERT_EQ(ascendcperf_v2::BroadcastPerf(node, perf), af::SUCCESS); | ||
| 122 | - ASSERT_GE(perf.ternary_ops.size(), 4U); | ||
| 123 | - for (const auto &entry : perf.ternary_ops) { | ||
| 124 | - EXPECT_FALSE(entry.second.GetTernaryOpStr().empty()); | ||
| 125 | - } | ||
| 126 | -} | ||
| 127 | - | ||
| 128 | -TEST(BroadcastNlastAxisPerfV2Symbolic, RankGreaterThanFourB8SkipsGatherLeaf) { | ||
| 129 | - auto node = MakeSymbolicNode({CreateExpr(2), CreateExpr(1), CreateExpr(2), CreateExpr(1), CreateExpr("last")}, | ||
| 130 | - {CreateExpr(2), CreateExpr(3), CreateExpr(2), CreateExpr(3), CreateExpr("last")}); | ||
| 131 | - node.input_dtype = {kUInt8}; | ||
| 132 | - node.output_dtype = {kUInt8}; | ||
| 133 | - EXPECT_EQ(ascendcperf_v2::GetNlastAxisPerfBranch(node), NlastAxisBranch::kDynamicLessThanVlUnaligned); | ||
| 134 | - EXPECT_FALSE(ascendcperf_v2::IsBroadcastFallback(node)); | ||
| 135 | -} | ||
| 136 | - | ||
| 137 | -} // namespace | ||
| 138 | -} // namespace att | ||
| @@ -1,88 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms and | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 of the License. | ||
| 5 | - * Please refer to the License for details. You should not use this file except in compliance with the License. | ||
| 6 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | - * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | - */ | ||
| 10 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | -namespace att { | ||
| 28 | -namespace { | ||
| 29 | -using ascendcapi_v2::VfCostAccumulator; | ||
| 30 | - | ||
| 31 | -TEST(BroadcastScalarDuplicatePerfV2, ScalarDuplicateCountsVectorRegisters) { | ||
| 32 | - const std::vector<std::tuple<std::string, int64_t, double>> cases = { | ||
| 33 | - {kFloat16, 8, 29}, {kFloat16, 129, 30}, {kFloat32, 65, 30}, {kFloat32, 45056, 732}, | ||
| 34 | - {kUInt8, 256, 29}, {kUInt8, 257, 30}, {kInt64, 33, 30}, {kInt64, 64, 30}, | ||
| 35 | - }; | ||
| 36 | - for (const auto &[dtype, duplicate_count, expected_value] : cases) { | ||
| 37 | - TensorShapeInfo input; | ||
| 38 | - input.data_type = dtype; | ||
| 39 | - TensorShapeInfo output; | ||
| 40 | - output.data_type = dtype; | ||
| 41 | - NodeInfo node; | ||
| 42 | - node.broadcast_node_params.valid = true; | ||
| 43 | - node.broadcast_node_params.is_scalar = true; | ||
| 44 | - node.broadcast_node_params.duplicate_count = CreateExpr(duplicate_count); | ||
| 45 | - | ||
| 46 | - const auto api = ApiPerfFactory::Instance().Create(kBroadcast + "V2"); | ||
| 47 | - ASSERT_NE(api, nullptr); | ||
| 48 | - PerfOutputInfo perf; | ||
| 49 | - ASSERT_EQ(api->GetPerfFunc()({input}, {output}, node, perf), af::SUCCESS); | ||
| 50 | - const auto iter = perf.pipe_res.find(PipeType::AIV_VEC); | ||
| 51 | - ASSERT_NE(iter, perf.pipe_res.end()); | ||
| 52 | - double value = 0.0; | ||
| 53 | - const std::vector<std::pair<af::Expression, af::Expression>> no_vars; | ||
| 54 | - ASSERT_EQ(iter->second.GetResult(no_vars, value), af::GRAPH_SUCCESS); | ||
| 55 | - EXPECT_EQ(value, expected_value) << "dtype=" << dtype << " duplicate_count=" << duplicate_count; | ||
| 56 | - } | ||
| 57 | -} | ||
| 58 | - | ||
| 59 | -TEST(BroadcastScalarDuplicatePerfV2, ScalarDuplicateSymbolicCountDividesByVectorLength) { | ||
| 60 | - TensorShapeInfo input; | ||
| 61 | - input.data_type = kFloat32; | ||
| 62 | - TensorShapeInfo output; | ||
| 63 | - output.data_type = kFloat32; | ||
| 64 | - NodeInfo node; | ||
| 65 | - node.broadcast_node_params.valid = true; | ||
| 66 | - node.broadcast_node_params.is_scalar = true; | ||
| 67 | - node.broadcast_node_params.duplicate_count = CreateExpr("z1z2t_size"); | ||
| 68 | - | ||
| 69 | - const auto api = ApiPerfFactory::Instance().Create(kBroadcast + "V2"); | ||
| 70 | - ASSERT_NE(api, nullptr); | ||
| 71 | - PerfOutputInfo perf; | ||
| 72 | - ASSERT_EQ(api->GetPerfFunc()({input}, {output}, node, perf), af::SUCCESS); | ||
| 73 | - const auto iter = perf.pipe_res.find(PipeType::AIV_VEC); | ||
| 74 | - ASSERT_NE(iter, perf.pipe_res.end()); | ||
| 75 | - | ||
| 76 | - Expr vl; | ||
| 77 | - Expr half_vl; | ||
| 78 | - Expr block_elements; | ||
| 79 | - ASSERT_EQ(ascendcapi_v2::GetBroadcastVectorElements(kFloat32, vl, half_vl, block_elements), af::SUCCESS); | ||
| 80 | - Expr repeat_time; | ||
| 81 | - ASSERT_EQ(ascendcapi_v2::CeilDiv(CreateExpr("z1z2t_size"), vl, repeat_time), af::SUCCESS); | ||
| 82 | - VfCostAccumulator acc; | ||
| 83 | - ASSERT_EQ(ascendcapi_v2::AddVfInstructPerf(kDuplicate, kFloat32, repeat_time, 1U, acc), af::SUCCESS); | ||
| 84 | - Expr expected = ascendcapi_v2::GetVfCost(acc); | ||
| 85 | - EXPECT_EQ(iter->second, expected); | ||
| 86 | -} | ||
| 87 | -} // namespace | ||
| 88 | -} // namespace att | ||
| @@ -12,7 +12,6 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | - | ||
| 16 | 15 | ||
| 17 | 16 | ||
| 18 | 17 | ||
| @@ -557,28 +556,26 @@ TEST(CodegenKernel, ConstantBrcStore) { | |||
| 557 | EXPECT_EQ(result, std::string{"Duplicate(local_1[0], scalar_0, local_1_actual_size);\n"}); | 556 | EXPECT_EQ(result, std::string{"Duplicate(local_1[0], scalar_0, local_1_actual_size);\n"}); |
| 558 | } | 557 | } |
| 559 | 558 | ||
| 560 | -TEST(CodegenKernel, BroadcastRegApiCall_BrcAlign_ABAWithTranspose) { | 559 | +TEST(CodegenKernel, BroadcastRegApiCall_BrcAlign_ABA) { |
| 561 | af::AscGraph graph("test_graph"); | 560 | af::AscGraph graph("test_graph"); |
| 562 | 561 | ||
| 563 | auto s0 = graph.CreateSizeVar(2); | 562 | auto s0 = graph.CreateSizeVar(2); |
| 564 | auto s1 = graph.CreateSizeVar(2); | 563 | auto s1 = graph.CreateSizeVar(2); |
| 565 | - auto s2 = graph.CreateSizeVar(8); | 564 | + auto s2 = graph.CreateSizeVar(7); |
| 566 | - auto z0 = graph.CreateAxis("z0", af::Axis::kAxisTypeTileInner, s0, {}, -1); | 565 | + auto z0 = graph.CreateAxis("z0", s0); |
| 567 | - auto z1 = graph.CreateAxis("z1", af::Axis::kAxisTypeTileInner, s1, {}, -1); | 566 | + auto z1 = graph.CreateAxis("z1", s1); |
| 568 | - auto z2 = graph.CreateAxis("z2", af::Axis::kAxisTypeTileInner, s2, {}, -1); | 567 | + auto z2 = graph.CreateAxis("z2", s2); |
| 569 | 568 | ||
| 570 | Data x_op("x", graph); | 569 | Data x_op("x", graph); |
| 571 | Load load_op("load"); | 570 | Load load_op("load"); |
| 572 | af::ascir_op::Broadcast broadcast_op("broadcast"); | 571 | af::ascir_op::Broadcast broadcast_op("broadcast"); |
| 573 | - af::ascir_op::Transpose transpose_op("transpose"); | ||
| 574 | graph.AddNode(load_op); | 572 | graph.AddNode(load_op); |
| 575 | graph.AddNode(broadcast_op); | 573 | graph.AddNode(broadcast_op); |
| 576 | - graph.AddNode(transpose_op); | ||
| 577 | 574 | ||
| 578 | load_op.x = x_op.y; | 575 | load_op.x = x_op.y; |
| 579 | load_op.attr.sched.axis = {z0.id, z1.id, z2.id}; | 576 | load_op.attr.sched.axis = {z0.id, z1.id, z2.id}; |
| 580 | *load_op.y.axis = {z0.id, z1.id, z2.id}; | 577 | *load_op.y.axis = {z0.id, z1.id, z2.id}; |
| 581 | - *load_op.y.repeats = {s0, One, s2}; // (2, 1, 8) | 578 | + *load_op.y.repeats = {s0, One, s2}; // (2, 1, 7) |
| 582 | *load_op.y.strides = {af::Symbol(8), Zero, One}; // (8, 0, 1) | 579 | *load_op.y.strides = {af::Symbol(8), Zero, One}; // (8, 0, 1) |
| 583 | broadcast_op.x = load_op.y; | 580 | broadcast_op.x = load_op.y; |
| 584 | *broadcast_op.y.axis = {z0.id, z1.id, z2.id}; | 581 | *broadcast_op.y.axis = {z0.id, z1.id, z2.id}; |
| @@ -636,14 +633,6 @@ TEST(CodegenKernel, BroadcastRegApiCall_BrcAlign_ABAWithTranspose) { | |||
| 636 | 633 | ||
| 637 | std::string result; | 634 | std::string result; |
| 638 | call.Generate(tpipe, vector<af::AxisId>{}, result); | 635 | call.Generate(tpipe, vector<af::AxisId>{}, result); |
| 639 | - auto params = ascir_param::GetAscirNodeParams(broadcast); | ||
| 640 | - ASSERT_NE(params, nullptr); | ||
| 641 | - const auto *broadcast_params = ascir_param::GetSpecificParams<ascir_param::BroadcastNodeParams>(*params); | ||
| 642 | - ASSERT_NE(broadcast_params, nullptr); | ||
| 643 | - ASSERT_EQ(broadcast_params->dst_shape.size(), 3UL); | ||
| 644 | - ASSERT_EQ(broadcast_params->src_shape.size(), 3UL); | ||
| 645 | - EXPECT_EQ(broadcast_params->dst_shape.back().role, ascir_param::ParamExprRole::kActualSize); | ||
| 646 | - EXPECT_EQ(broadcast_params->src_shape.back().role, ascir_param::ParamExprRole::kActualSize); | ||
| 647 | EXPECT_EQ(result, | 636 | EXPECT_EQ(result, |
| 648 | "const uint32_t dst_shape_0_brc_to_1[3] = {static_cast<uint32_t>(2), static_cast<uint32_t>(2), " | 637 | "const uint32_t dst_shape_0_brc_to_1[3] = {static_cast<uint32_t>(2), static_cast<uint32_t>(2), " |
| 649 | "static_cast<uint32_t>(8)};\n" | 638 | "static_cast<uint32_t>(8)};\n" |
| @@ -1,178 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | -namespace att { | ||
| 17 | -namespace ascendcperf_v2 { | ||
| 18 | -namespace { | ||
| 19 | -using ascendcapi_v2::BroadcastTilingInfo; | ||
| 20 | -using ascendcapi_v2::ParamExprInputs; | ||
| 21 | -using ascendcapi_v2::VfCostAccumulator; | ||
| 22 | - | ||
| 23 | -ParamExprInputs MakeParamInputs(const std::vector<Expr> &dims, const std::vector<Expr> &actual_dims) { | ||
| 24 | - ParamExprInputs inputs; | ||
| 25 | - inputs.semantic = dims; | ||
| 26 | - inputs.size = dims; | ||
| 27 | - inputs.actual_size = actual_dims.size() == dims.size() ? actual_dims : dims; | ||
| 28 | - return inputs; | ||
| 29 | -} | ||
| 30 | - | ||
| 31 | -af::Status SetBroadcastCost(const VfCostAccumulator &acc, PerfOutputInfo &perf) { | ||
| 32 | - perf.pipe_res[PipeType::AIV_VEC] = ascendcapi_v2::GetVfCost(acc); | ||
| 33 | - return af::SUCCESS; | ||
| 34 | -} | ||
| 35 | - | ||
| 36 | -af::Status BuildScalarPerf(const NodeDetail &node_info, PerfOutputInfo &perf) { | ||
| 37 | - GE_ASSERT_TRUE(!node_info.input_dtype.empty(), "Broadcast scalar input dtype is missing."); | ||
| 38 | - std::string helper_dtype; | ||
| 39 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetEffectiveHelperDtype(node_info.input_dtype[0], helper_dtype)); | ||
| 40 | - Expr vl; | ||
| 41 | - Expr half_vl; | ||
| 42 | - Expr block_elements; | ||
| 43 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetBroadcastVectorElements(node_info.input_dtype[0], vl, half_vl, block_elements)); | ||
| 44 | - (void)half_vl; | ||
| 45 | - (void)block_elements; | ||
| 46 | - Expr repeat_time; | ||
| 47 | - GE_ASSERT_SUCCESS(ascendcapi_v2::CeilDiv(node_info.broadcast_node_params.duplicate_count, vl, repeat_time)); | ||
| 48 | - VfCostAccumulator acc; | ||
| 49 | - GE_ASSERT_SUCCESS(ascendcapi_v2::AddVfInstructPerf(kDuplicate, helper_dtype, repeat_time, 1U, acc)); | ||
| 50 | - return SetBroadcastCost(acc, perf); | ||
| 51 | -} | ||
| 52 | - | ||
| 53 | -af::Status BuildEqualSizePerf(const NodeDetail &node_info, PerfOutputInfo &perf) { | ||
| 54 | - GE_ASSERT_TRUE(!node_info.input_dtype.empty(), "Broadcast equal-size input dtype is missing."); | ||
| 55 | - VfCostAccumulator acc; | ||
| 56 | - GE_ASSERT_SUCCESS(ascendcapi_v2::AddBroadcastDataCopyPerf(node_info.input_dtype[0], CreateExpr(1), acc)); | ||
| 57 | - return SetBroadcastCost(acc, perf); | ||
| 58 | -} | ||
| 59 | - | ||
| 60 | -af::Status BuildSingleElementPerf(const NodeDetail &node_info, const BroadcastTilingInfo &tiling, | ||
| 61 | - PerfOutputInfo &perf) { | ||
| 62 | - GE_ASSERT_TRUE(!node_info.input_dtype.empty(), "Broadcast duplicate input dtype is missing."); | ||
| 63 | - std::string helper_dtype; | ||
| 64 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetEffectiveHelperDtype(node_info.input_dtype[0], helper_dtype)); | ||
| 65 | - Expr vl; | ||
| 66 | - Expr half_vl; | ||
| 67 | - Expr block_elements; | ||
| 68 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetBroadcastVectorElements(node_info.input_dtype[0], vl, half_vl, block_elements)); | ||
| 69 | - (void)half_vl; | ||
| 70 | - (void)block_elements; | ||
| 71 | - Expr repeat_time; | ||
| 72 | - GE_ASSERT_SUCCESS(ascendcapi_v2::CeilDiv(tiling.dst_size, vl, repeat_time)); | ||
| 73 | - VfCostAccumulator acc; | ||
| 74 | - GE_ASSERT_SUCCESS(ascendcapi_v2::AddVfInstructPerf(kLoad, node_info.input_dtype[0], CreateExpr(1), 1U, acc)); | ||
| 75 | - GE_ASSERT_SUCCESS(ascendcapi_v2::AddVfInstructPerf(kUpdateMask, helper_dtype, repeat_time, 1U, acc)); | ||
| 76 | - GE_ASSERT_SUCCESS(ascendcapi_v2::AddVfInstructPerf(kStore, node_info.output_dtype[0], repeat_time, 1U, acc)); | ||
| 77 | - return SetBroadcastCost(acc, perf); | ||
| 78 | -} | ||
| 79 | - | ||
| 80 | -af::Status ValidateBroadcastNodeDetail(const NodeDetail &node_info) { | ||
| 81 | - const auto ¶ms = node_info.broadcast_node_params; | ||
| 82 | - GE_ASSERT_TRUE(!node_info.input_dtype.empty() && !node_info.output_dtype.empty(), | ||
| 83 | - "Broadcast dtype information is missing."); | ||
| 84 | - if (params.is_scalar) { | ||
| 85 | - return af::SUCCESS; | ||
| 86 | - } | ||
| 87 | - GE_ASSERT_TRUE(!node_info.input_dims.empty() && node_info.input_dims.size() == node_info.output_dims.size(), | ||
| 88 | - "Broadcast shape ranks are invalid."); | ||
| 89 | - GE_ASSERT_TRUE( | ||
| 90 | - params.src_shape.size() == params.dst_shape.size() && params.src_shape.size() == node_info.input_dims.size(), | ||
| 91 | - "Broadcast parameter shape rank is invalid."); | ||
| 92 | - GE_ASSERT_TRUE(node_info.repeats.size() == node_info.input_dims.size(), "Broadcast input repeats length is invalid."); | ||
| 93 | - return af::SUCCESS; | ||
| 94 | -} | ||
| 95 | - | ||
| 96 | -Expr ResolveBroadcastSize(const std::vector<ascir_param::ParamExprLeaf> &shape, const ParamExprInputs &inputs) { | ||
| 97 | - std::vector<Expr> resolved; | ||
| 98 | - for (size_t i = 0U; i < shape.size(); ++i) { | ||
| 99 | - resolved.push_back( | ||
| 100 | - ascendcapi_v2::ResolveParamExprLeaf(shape[i], inputs.semantic[i], inputs.size[i], inputs.actual_size[i])); | ||
| 101 | - } | ||
| 102 | - return ascendcapi_v2::ShapeProduct(resolved); | ||
| 103 | -} | ||
| 104 | -} // namespace | ||
| 105 | - | ||
| 106 | -bool IsBroadcastFallback(const NodeDetail &node_info) { | ||
| 107 | - if (node_info.input_dtype.empty() || node_info.input_dims.empty() || node_info.output_dims.empty()) { | ||
| 108 | - return false; | ||
| 109 | - } | ||
| 110 | - const auto ¶ms = node_info.broadcast_node_params; | ||
| 111 | - const auto src_inputs = MakeParamInputs(node_info.input_dims, node_info.repeats); | ||
| 112 | - const auto dst_inputs = MakeParamInputs(node_info.output_dims, node_info.output_dims); | ||
| 113 | - BroadcastTilingInfo tiling; | ||
| 114 | - if (BuildBroadcastTiling(params.src_shape, params.dst_shape, node_info.input_dtype[0], src_inputs, dst_inputs, | ||
| 115 | - tiling) != af::SUCCESS) { | ||
| 116 | - return false; | ||
| 117 | - } | ||
| 118 | - const Expr src_size = ResolveBroadcastSize(params.src_shape, src_inputs); | ||
| 119 | - const Expr dst_size = ResolveBroadcastSize(params.dst_shape, dst_inputs); | ||
| 120 | - if (src_size == dst_size || src_size == CreateExpr(1)) { | ||
| 121 | - return false; | ||
| 122 | - } | ||
| 123 | - if (IsLastAxisBroadcast(tiling)) { | ||
| 124 | - return GetLastAxisBranch(tiling, node_info.input_dtype[0], params.const_rank) == | ||
| 125 | - ascendcapi_v2::LastAxisBranch::kFallback; | ||
| 126 | - } | ||
| 127 | - if (params.const_rank == -1 && ascendcperf_v2::HasUnknownNlastCondition(node_info, tiling)) { | ||
| 128 | - return false; | ||
| 129 | - } | ||
| 130 | - if (tiling.original_rank > 4U) { | ||
| 131 | - return false; | ||
| 132 | - } | ||
| 133 | - return GetNlastAxisBranch(tiling, node_info.input_dtype[0], params.const_rank) == | ||
| 134 | - ascendcapi_v2::NlastAxisBranch::kFallback; | ||
| 135 | -} | ||
| 136 | - | ||
| 137 | -af::Status BroadcastPerf(const NodeDetail &node_info, PerfOutputInfo &perf) { | ||
| 138 | - GE_ASSERT_TRUE(node_info.broadcast_node_params.valid, "Broadcast parameters are invalid."); | ||
| 139 | - GE_ASSERT_SUCCESS(ValidateBroadcastNodeDetail(node_info)); | ||
| 140 | - if (node_info.broadcast_node_params.is_scalar) { | ||
| 141 | - return BuildScalarPerf(node_info, perf); | ||
| 142 | - } | ||
| 143 | - GE_ASSERT_TRUE(!node_info.input_dtype.empty() && !node_info.output_dtype.empty(), | ||
| 144 | - "Broadcast dtype information is missing."); | ||
| 145 | - const auto ¶ms = node_info.broadcast_node_params; | ||
| 146 | - const auto src_inputs = MakeParamInputs(node_info.input_dims, node_info.repeats); | ||
| 147 | - const auto dst_inputs = MakeParamInputs(node_info.output_dims, node_info.output_dims); | ||
| 148 | - const Expr raw_src_size = ResolveBroadcastSize(params.src_shape, src_inputs); | ||
| 149 | - const Expr raw_dst_size = ResolveBroadcastSize(params.dst_shape, dst_inputs); | ||
| 150 | - BroadcastTilingInfo tiling; | ||
| 151 | - GE_ASSERT_SUCCESS(ascendcapi_v2::BuildBroadcastTiling(params.src_shape, params.dst_shape, node_info.input_dtype[0], | ||
| 152 | - src_inputs, dst_inputs, tiling)); | ||
| 153 | - if (raw_src_size == raw_dst_size) { | ||
| 154 | - return BuildEqualSizePerf(node_info, perf); | ||
| 155 | - } | ||
| 156 | - if (raw_src_size == CreateExpr(1)) { | ||
| 157 | - tiling.src_size = raw_src_size; | ||
| 158 | - tiling.dst_size = raw_dst_size; | ||
| 159 | - return BuildSingleElementPerf(node_info, tiling, perf); | ||
| 160 | - } | ||
| 161 | - if (ascendcapi_v2::IsLastAxisBroadcast(tiling)) { | ||
| 162 | - const auto branch = | ||
| 163 | - ascendcapi_v2::GetLastAxisBranch(tiling, node_info.input_dtype[0], node_info.broadcast_node_params.const_rank); | ||
| 164 | - if (branch == ascendcapi_v2::LastAxisBranch::kFallback && params.const_rank != -1) { | ||
| 165 | - return af::FAILED; | ||
| 166 | - } | ||
| 167 | - return BuildLastAxisPerf(node_info, tiling, perf); | ||
| 168 | - } | ||
| 169 | - const auto branch = | ||
| 170 | - ascendcapi_v2::GetNlastAxisBranch(tiling, node_info.input_dtype[0], node_info.broadcast_node_params.const_rank); | ||
| 171 | - if (branch == ascendcapi_v2::NlastAxisBranch::kFallback && params.const_rank != -1 && tiling.original_rank <= 4U) { | ||
| 172 | - return af::FAILED; | ||
| 173 | - } | ||
| 174 | - return BuildNlastAxisPerf(node_info, tiling, perf); | ||
| 175 | -} | ||
| 176 | - | ||
| 177 | -} // namespace ascendcperf_v2 | ||
| 178 | -} // namespace att | ||
| @@ -1,21 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | -namespace att { | ||
| 13 | -namespace ascendcperf_v2 { | ||
| 14 | - | ||
| 15 | -af::Status BroadcastPerf(const NodeDetail &node_info, PerfOutputInfo &perf); | ||
| 16 | -bool IsBroadcastFallback(const NodeDetail &node_info); | ||
| 17 | - | ||
| 18 | -} // namespace ascendcperf_v2 | ||
| 19 | -} // namespace att | ||
| 20 | - | ||
| 21 | - | ||
| @@ -1,669 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 of the License. | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | -namespace att { | ||
| 13 | -namespace ascendcapi_v2 { | ||
| 14 | -namespace { | ||
| 15 | - | ||
| 16 | -bool IsTrue(const Expr &lhs, CondType condition, const Expr &rhs) { | ||
| 17 | - af::TriBool result = af::TriBool::kUnknown; | ||
| 18 | - if (condition == CondType::K_EQ) { | ||
| 19 | - result = af::SymbolicUtils::StaticCheckEq(lhs, rhs); | ||
| 20 | - } | ||
| 21 | - if (condition == CondType::K_LT) { | ||
| 22 | - result = af::SymbolicUtils::StaticCheckLt(lhs, rhs); | ||
| 23 | - } | ||
| 24 | - if (condition == CondType::K_LE) { | ||
| 25 | - result = af::SymbolicUtils::StaticCheckLe(lhs, rhs); | ||
| 26 | - } | ||
| 27 | - if (condition == CondType::K_GT) { | ||
| 28 | - result = af::SymbolicUtils::StaticCheckGt(lhs, rhs); | ||
| 29 | - } | ||
| 30 | - return result == af::TriBool::kTrue; | ||
| 31 | -} | ||
| 32 | - | ||
| 33 | -bool IsB8(const std::string &dtype) { | ||
| 34 | - return dtype == kUInt8 || dtype == kInt8; | ||
| 35 | -} | ||
| 36 | - | ||
| 37 | -bool IsZero(const Expr &value) { | ||
| 38 | - return af::SymbolicUtils::StaticCheckEq(value, CreateExpr(0)) == af::TriBool::kTrue; | ||
| 39 | -} | ||
| 40 | - | ||
| 41 | -LastAxisBranch ClassifyRankTwo(const BroadcastTilingInfo &tiling, const std::string &dtype, const Expr &vl, | ||
| 42 | - const Expr &half_vl, const Expr &block, int32_t const_rank) { | ||
| 43 | - const Expr last = tiling.dst_shape[1U]; | ||
| 44 | - if (IsTrue(last, CondType::K_EQ, block) && !IsB8(dtype)) { | ||
| 45 | - return LastAxisBranch::kE2B; | ||
| 46 | - } | ||
| 47 | - if (IsTrue(last, CondType::K_LT, half_vl)) { | ||
| 48 | - return IsTrue(tiling.dst_size, CondType::K_LT, vl) ? LastAxisBranch::kGatherOne : LastAxisBranch::kGatherTwo; | ||
| 49 | - } | ||
| 50 | - if (IsTrue(last, CondType::K_LE, vl)) { | ||
| 51 | - return LastAxisBranch::kLessThanVlUnaligned; | ||
| 52 | - } | ||
| 53 | - if (IsTrue(af::sym::Mod(last, block), CondType::K_EQ, CreateExpr(0))) { | ||
| 54 | - return LastAxisBranch::kLargerThanVlAligned; | ||
| 55 | - } | ||
| 56 | - return const_rank == -1 ? LastAxisBranch::kDynamicLargerThanVlUnaligned : LastAxisBranch::kLargerThanVlUnaligned; | ||
| 57 | -} | ||
| 58 | - | ||
| 59 | -LastAxisBranch ClassifyRankThree(const BroadcastTilingInfo &tiling, const std::string &dtype, const Expr &vl, | ||
| 60 | - const Expr &half_vl, const Expr &block, int32_t const_rank) { | ||
| 61 | - const Expr last = tiling.dst_shape[2U]; | ||
| 62 | - if (IsZero(block)) { | ||
| 63 | - return LastAxisBranch::kFallback; | ||
| 64 | - } | ||
| 65 | - const Expr default_block_num = vl / block; | ||
| 66 | - if (!IsB8(dtype) && tiling.src_stride.size() > 1U && !IsZero(tiling.src_stride[1]) && | ||
| 67 | - IsTrue(last, CondType::K_EQ, block) && IsTrue(tiling.dst_shape[1] * last, CondType::K_GT, half_vl) && | ||
| 68 | - IsTrue(af::sym::Mod(tiling.dst_shape[1], default_block_num), CondType::K_EQ, CreateExpr(0))) { | ||
| 69 | - return IsTrue(tiling.dst_shape[1] * last, CondType::K_GT, vl) ? LastAxisBranch::kE2BLargerThanVl | ||
| 70 | - : LastAxisBranch::kE2BLessThanVl; | ||
| 71 | - } | ||
| 72 | - if (IsTrue(last, CondType::K_LT, half_vl)) { | ||
| 73 | - return LastAxisBranch::kGatherWrapper; | ||
| 74 | - } | ||
| 75 | - if (IsTrue(last, CondType::K_LE, vl)) { | ||
| 76 | - if (IsTrue(af::sym::Mod(last, block), CondType::K_EQ, CreateExpr(0))) { | ||
| 77 | - return LastAxisBranch::kLessThanVlAligned; | ||
| 78 | - } | ||
| 79 | - return const_rank == -1 ? LastAxisBranch::kDynamicLessThanVlUnaligned : LastAxisBranch::kLessThanVlUnaligned; | ||
| 80 | - } | ||
| 81 | - if (IsTrue(af::sym::Mod(last, block), CondType::K_EQ, CreateExpr(0))) { | ||
| 82 | - return LastAxisBranch::kLargerThanVlAligned; | ||
| 83 | - } | ||
| 84 | - return const_rank == -1 ? LastAxisBranch::kDynamicLargerThanVlUnaligned : LastAxisBranch::kLargerThanVlUnaligned; | ||
| 85 | -} | ||
| 86 | - | ||
| 87 | -LastAxisBranch ClassifyRankFour(const BroadcastTilingInfo &tiling, const std::string &dtype, const Expr &vl, | ||
| 88 | - const Expr &half_vl, const Expr &block, int32_t const_rank) { | ||
| 89 | - const Expr last = tiling.dst_shape[3U]; | ||
| 90 | - if (IsZero(block)) { | ||
| 91 | - return LastAxisBranch::kFallback; | ||
| 92 | - } | ||
| 93 | - const Expr default_block_num = vl / block; | ||
| 94 | - if (!IsB8(dtype) && tiling.src_stride.size() > 2U && !IsZero(tiling.src_stride[2]) && | ||
| 95 | - IsTrue(last, CondType::K_EQ, block) && | ||
| 96 | - IsTrue(af::sym::Mod(tiling.dst_shape[2], default_block_num), CondType::K_EQ, CreateExpr(0))) { | ||
| 97 | - return LastAxisBranch::kE2B; | ||
| 98 | - } | ||
| 99 | - if (IsTrue(last, CondType::K_LT, half_vl) && !IsB8(dtype)) { | ||
| 100 | - return LastAxisBranch::kGatherWrapperForFourDim; | ||
| 101 | - } | ||
| 102 | - if (IsTrue(last, CondType::K_LE, vl)) { | ||
| 103 | - if (IsTrue(af::sym::Mod(last, block), CondType::K_EQ, CreateExpr(0))) { | ||
| 104 | - return LastAxisBranch::kLessThanVlAligned; | ||
| 105 | - } | ||
| 106 | - return const_rank == -1 ? LastAxisBranch::kDynamicLessThanVlUnaligned : LastAxisBranch::kLessThanVlUnaligned; | ||
| 107 | - } | ||
| 108 | - if (IsTrue(af::sym::Mod(last, block), CondType::K_EQ, CreateExpr(0))) { | ||
| 109 | - return LastAxisBranch::kLargerThanVlAligned; | ||
| 110 | - } | ||
| 111 | - return const_rank == -1 ? LastAxisBranch::kDynamicLargerThanVlUnaligned : LastAxisBranch::kLargerThanVlUnaligned; | ||
| 112 | -} | ||
| 113 | - | ||
| 114 | -LastAxisBranch ClassifySmallRank(const BroadcastTilingInfo &tiling, const std::string &dtype, int32_t const_rank) { | ||
| 115 | - Expr vl; | ||
| 116 | - Expr half_vl; | ||
| 117 | - Expr block; | ||
| 118 | - if (GetBroadcastTilingVectorElements(dtype, vl, half_vl, block) != af::SUCCESS) { | ||
| 119 | - return LastAxisBranch::kFallback; | ||
| 120 | - } | ||
| 121 | - if (tiling.rank == 1U) { | ||
| 122 | - return LastAxisBranch::kFallback; | ||
| 123 | - } | ||
| 124 | - if (tiling.rank == 2U) { | ||
| 125 | - return ClassifyRankTwo(tiling, dtype, vl, half_vl, block, const_rank); | ||
| 126 | - } | ||
| 127 | - if (tiling.rank == 3U) { | ||
| 128 | - return ClassifyRankThree(tiling, dtype, vl, half_vl, block, const_rank); | ||
| 129 | - } | ||
| 130 | - return ClassifyRankFour(tiling, dtype, vl, half_vl, block, const_rank); | ||
| 131 | -} | ||
| 132 | - | ||
| 133 | -} // namespace | ||
| 134 | - | ||
| 135 | -LastAxisBranch GetLastAxisBranch(const BroadcastTilingInfo &tiling, const std::string &dtype, int32_t const_rank) { | ||
| 136 | - if (tiling.rank == 0U || tiling.dst_shape.size() < tiling.rank) { | ||
| 137 | - return LastAxisBranch::kFallback; | ||
| 138 | - } | ||
| 139 | - if (tiling.rank > 4U) { | ||
| 140 | - Expr vl; | ||
| 141 | - Expr half_vl; | ||
| 142 | - Expr block; | ||
| 143 | - if (GetBroadcastTilingVectorElements(dtype, vl, half_vl, block) != af::SUCCESS) { | ||
| 144 | - return LastAxisBranch::kFallback; | ||
| 145 | - } | ||
| 146 | - const Expr last = tiling.dst_shape.back(); | ||
| 147 | - const auto last_check = af::SymbolicUtils::StaticCheckLe(last, vl); | ||
| 148 | - if (last_check == af::TriBool::kTrue) { | ||
| 149 | - return LastAxisBranch::kDynamicLessThanVlUnaligned; | ||
| 150 | - } | ||
| 151 | - if (last_check == af::TriBool::kFalse) { | ||
| 152 | - return LastAxisBranch::kDynamicLargerThanVlUnaligned; | ||
| 153 | - } | ||
| 154 | - return LastAxisBranch::kFallback; | ||
| 155 | - } | ||
| 156 | - return ClassifySmallRank(tiling, dtype, tiling.original_rank > 4U ? -1 : const_rank); | ||
| 157 | -} | ||
| 158 | - | ||
| 159 | -} // namespace ascendcapi_v2 | ||
| 160 | - | ||
| 161 | -namespace ascendcperf_v2 { | ||
| 162 | -namespace { | ||
| 163 | -using ascendcapi_v2::BroadcastTilingInfo; | ||
| 164 | -using ascendcapi_v2::LastAxisBranch; | ||
| 165 | -using ascendcapi_v2::ParamExprInputs; | ||
| 166 | -using ascendcapi_v2::VfCostAccumulator; | ||
| 167 | - | ||
| 168 | -bool IsB8Dtype(const std::string &dtype) { | ||
| 169 | - return dtype == kUInt8 || dtype == kInt8; | ||
| 170 | -} | ||
| 171 | - | ||
| 172 | -std::string HelperDtype(const std::string &dtype) { | ||
| 173 | - std::string helper_dtype; | ||
| 174 | - return ascendcapi_v2::GetEffectiveHelperDtype(dtype, helper_dtype) == af::SUCCESS ? helper_dtype : dtype; | ||
| 175 | -} | ||
| 176 | - | ||
| 177 | -Expr ProductFrom(const std::vector<Expr> &shape, size_t end) { | ||
| 178 | - Expr result = CreateExpr(1); | ||
| 179 | - for (size_t i = 0U; i < end; ++i) { | ||
| 180 | - result = result * shape[i]; | ||
| 181 | - } | ||
| 182 | - return result; | ||
| 183 | -} | ||
| 184 | - | ||
| 185 | -af::Status Add(const std::string &op, const std::string &dtype, const Expr &count, VfCostAccumulator &acc) { | ||
| 186 | - return ascendcapi_v2::AddVfInstructPerf(op, dtype, count, 1U, acc); | ||
| 187 | -} | ||
| 188 | - | ||
| 189 | -af::Status AddLastGatherIndexPerf(const std::string &dtype, VfCostAccumulator &acc) { | ||
| 190 | - Expr byte_size; | ||
| 191 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetDtypeByteSize(dtype, byte_size)); | ||
| 192 | - const std::string index_dtype = | ||
| 193 | - af::SymbolicUtils::StaticCheckLe(byte_size, CreateExpr(2)) == af::TriBool::kTrue ? kInt16 : kInt32; | ||
| 194 | - const Expr calls = IsB8Dtype(dtype) ? CreateExpr(2) : CreateExpr(1); | ||
| 195 | - GE_ASSERT_SUCCESS(Add(kDuplicate, index_dtype, calls, acc)); | ||
| 196 | - // Placeholder: Reg::Arange in last-axis index generation, index dtype, one index batch. | ||
| 197 | - GE_ASSERT_SUCCESS(VfPerfUtils::AddVfInstructPerf(kPlaceholder, index_dtype, acc.max_latency, acc.throughput, calls)); | ||
| 198 | - GE_ASSERT_SUCCESS(Add(kDiv, index_dtype, calls, acc)); | ||
| 199 | - return Add(kStore, index_dtype, calls, acc); | ||
| 200 | -} | ||
| 201 | - | ||
| 202 | -bool IsLt(const Expr &lhs, const Expr &rhs) { | ||
| 203 | - return af::SymbolicUtils::StaticCheckLt(lhs, rhs) == af::TriBool::kTrue; | ||
| 204 | -} | ||
| 205 | - | ||
| 206 | -af::Status SafeDiv(const Expr &value, const Expr &divisor, Expr &result) { | ||
| 207 | - GE_ASSERT_TRUE(af::SymbolicUtils::StaticCheckEq(divisor, CreateExpr(0)) != af::TriBool::kTrue, | ||
| 208 | - "Broadcast last-axis divisor cannot be zero."); | ||
| 209 | - result = af::sym::Floor(value / divisor); | ||
| 210 | - return af::SUCCESS; | ||
| 211 | -} | ||
| 212 | - | ||
| 213 | -std::string GatherIndexDtype(const std::string &dtype) { | ||
| 214 | - std::string helper_dtype; | ||
| 215 | - if (ascendcapi_v2::GetEffectiveHelperDtype(dtype, helper_dtype) == af::SUCCESS && helper_dtype != dtype) { | ||
| 216 | - return GatherIndexDtype(helper_dtype); | ||
| 217 | - } | ||
| 218 | - Expr bytes; | ||
| 219 | - if (ascendcapi_v2::GetDtypeByteSize(dtype, bytes) != af::SUCCESS) { | ||
| 220 | - return dtype; | ||
| 221 | - } | ||
| 222 | - return bytes == kSymFour ? kInt32 : kInt16; | ||
| 223 | -} | ||
| 224 | - | ||
| 225 | -std::string GatherRegDtype(const std::string &dtype) { | ||
| 226 | - std::string helper_dtype; | ||
| 227 | - if (ascendcapi_v2::GetEffectiveHelperDtype(dtype, helper_dtype) == af::SUCCESS && helper_dtype != dtype) { | ||
| 228 | - return GatherRegDtype(helper_dtype); | ||
| 229 | - } | ||
| 230 | - Expr bytes; | ||
| 231 | - if (ascendcapi_v2::GetDtypeByteSize(dtype, bytes) != af::SUCCESS) { | ||
| 232 | - return dtype; | ||
| 233 | - } | ||
| 234 | - return bytes == kSymFour ? kUInt32 : kUInt16; | ||
| 235 | -} | ||
| 236 | - | ||
| 237 | -std::string RankTwoGatherDtype(const std::string &dtype) { | ||
| 238 | - return GatherRegDtype(dtype); | ||
| 239 | -} | ||
| 240 | - | ||
| 241 | -af::Status AddWrapperIndexGeneration(const std::string &dtype, size_t rank, VfCostAccumulator &acc) { | ||
| 242 | - const std::string index_dtype = GatherIndexDtype(dtype); | ||
| 243 | - const Expr lanes = IsB8Dtype(dtype) ? CreateExpr(2) : CreateExpr(1); | ||
| 244 | - GE_ASSERT_SUCCESS(Add(kDuplicate, index_dtype, CreateExpr(static_cast<int64_t>(rank * 2U)), acc)); | ||
| 245 | - // Placeholder: Reg::Arange in VfGenIndex/VfGenIndexB8. | ||
| 246 | - GE_ASSERT_SUCCESS(VfPerfUtils::AddVfInstructPerf(kPlaceholder, index_dtype, acc.max_latency, acc.throughput, lanes)); | ||
| 247 | - GE_ASSERT_SUCCESS(Add(kDiv, index_dtype, CreateExpr(static_cast<int64_t>(rank)) * lanes, acc)); | ||
| 248 | - GE_ASSERT_SUCCESS(Add(kMul, index_dtype, CreateExpr(static_cast<int64_t>(rank + 1U)) * lanes, acc)); | ||
| 249 | - GE_ASSERT_SUCCESS(Add(kSub, index_dtype, CreateExpr(static_cast<int64_t>(rank)) * lanes, acc)); | ||
| 250 | - GE_ASSERT_SUCCESS(Add(kMulAddDst, index_dtype, CreateExpr(static_cast<int64_t>(rank - 1U)) * lanes, acc)); | ||
| 251 | - return Add(kStore, index_dtype, lanes, acc); | ||
| 252 | -} | ||
| 253 | - | ||
| 254 | -struct GatherLoopCounts { | ||
| 255 | - Expr first{CreateExpr(1)}; | ||
| 256 | - Expr second{CreateExpr(1)}; | ||
| 257 | - Expr third{CreateExpr(1)}; | ||
| 258 | -}; | ||
| 259 | - | ||
| 260 | -af::Status GetRankThreeGatherLoops(const std::vector<Expr> &shape, const Expr &vl, GatherLoopCounts &counts) { | ||
| 261 | - const Expr inner = shape[2] * shape[1]; | ||
| 262 | - Expr tile; | ||
| 263 | - if (IsLt(inner, vl)) { | ||
| 264 | - GE_ASSERT_SUCCESS(SafeDiv(vl, inner, tile)); | ||
| 265 | - return SafeDiv(shape[0], tile, counts.second); | ||
| 266 | - } | ||
| 267 | - GE_ASSERT_SUCCESS(SafeDiv(vl, shape[2], tile)); | ||
| 268 | - counts.first = shape[0]; | ||
| 269 | - return SafeDiv(shape[1], tile, counts.second); | ||
| 270 | -} | ||
| 271 | - | ||
| 272 | -af::Status GetRankFourGatherLoops(const std::vector<Expr> &shape, const Expr &vl, GatherLoopCounts &counts) { | ||
| 273 | - const Expr inner_three = shape[3] * shape[2] * shape[1]; | ||
| 274 | - Expr tile; | ||
| 275 | - if (IsLt(inner_three, vl)) { | ||
| 276 | - GE_ASSERT_SUCCESS(SafeDiv(vl, inner_three, tile)); | ||
| 277 | - return SafeDiv(shape[0], tile, counts.third); | ||
| 278 | - } | ||
| 279 | - const Expr inner_two = shape[3] * shape[2]; | ||
| 280 | - if (IsLt(inner_two, vl)) { | ||
| 281 | - GE_ASSERT_SUCCESS(SafeDiv(vl, inner_two, tile)); | ||
| 282 | - counts.second = shape[0]; | ||
| 283 | - return SafeDiv(shape[1], tile, counts.third); | ||
| 284 | - } | ||
| 285 | - GE_ASSERT_SUCCESS(SafeDiv(vl, shape[3], tile)); | ||
| 286 | - counts.first = shape[0]; | ||
| 287 | - counts.second = shape[1]; | ||
| 288 | - return SafeDiv(shape[2], tile, counts.third); | ||
| 289 | -} | ||
| 290 | - | ||
| 291 | -af::Status AddGatherWrapperArithmetic(const std::string &dtype, size_t rank, const GatherLoopCounts &loops, | ||
| 292 | - VfCostAccumulator &acc) { | ||
| 293 | - const Expr lanes = IsB8Dtype(dtype) ? CreateExpr(2) : CreateExpr(1); | ||
| 294 | - const Expr first = loops.first; | ||
| 295 | - const Expr second = loops.first * (rank == 3U ? loops.second + CreateExpr(1) : loops.second); | ||
| 296 | - const Expr third = rank == 3U ? CreateExpr(0) : loops.first * loops.second * (loops.third + CreateExpr(1)); | ||
| 297 | - const Expr arithmetic = first + second + third; | ||
| 298 | - GE_ASSERT_SUCCESS(Add(kMuls, GatherRegDtype(dtype), arithmetic, acc)); | ||
| 299 | - return Add(kAdd, GatherRegDtype(dtype), arithmetic * lanes, acc); | ||
| 300 | -} | ||
| 301 | - | ||
| 302 | -af::Status AddGatherWrapperPerf(const NodeDetail &node, const BroadcastTilingInfo &tiling, size_t rank, | ||
| 303 | - VfCostAccumulator &acc) { | ||
| 304 | - Expr vl; | ||
| 305 | - Expr half_vl; | ||
| 306 | - Expr block; | ||
| 307 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetBroadcastTilingVectorElements(node.input_dtype[0], vl, half_vl, block)); | ||
| 308 | - (void)half_vl; | ||
| 309 | - (void)block; | ||
| 310 | - GatherLoopCounts loops; | ||
| 311 | - const auto &shape = tiling.dst_shape; | ||
| 312 | - GE_ASSERT_SUCCESS(rank == 3U ? GetRankThreeGatherLoops(shape, vl, loops) : GetRankFourGatherLoops(shape, vl, loops)); | ||
| 313 | - const Expr calls = rank == 3U ? loops.first * (loops.second + CreateExpr(1)) | ||
| 314 | - : loops.first * loops.second * (loops.third + CreateExpr(1)); | ||
| 315 | - GE_ASSERT_SUCCESS(AddWrapperIndexGeneration(node.input_dtype[0], rank, acc)); | ||
| 316 | - GE_ASSERT_SUCCESS(Add(kDuplicate, GatherIndexDtype(node.input_dtype[0]), kSymTwo, acc)); | ||
| 317 | - const Expr lanes = IsB8Dtype(node.input_dtype[0]) ? CreateExpr(2) : CreateExpr(1); | ||
| 318 | - GE_ASSERT_SUCCESS(Add(kLoad, GatherIndexDtype(node.input_dtype[0]), lanes, acc)); | ||
| 319 | - GE_ASSERT_SUCCESS(AddGatherWrapperArithmetic(node.input_dtype[0], rank, loops, acc)); | ||
| 320 | - const Expr gathers = IsB8Dtype(node.input_dtype[0]) ? calls * CreateExpr(2) : calls; | ||
| 321 | - GE_ASSERT_SUCCESS(VfPerfUtils::AddVfInstructPerf(kPlaceholder, GatherRegDtype(node.input_dtype[0]), acc.max_latency, | ||
| 322 | - acc.throughput, gathers)); | ||
| 323 | - if (IsB8Dtype(node.input_dtype[0])) { | ||
| 324 | - GE_ASSERT_SUCCESS(Add(kDeInterleave, node.input_dtype[0], calls, acc)); | ||
| 325 | - } | ||
| 326 | - GE_ASSERT_SUCCESS(Add(kStore, node.output_dtype[0], calls, acc)); | ||
| 327 | - return VfPerfUtils::AddVfInstructPerf(kPlaceholder, node.output_dtype[0], acc.max_latency, acc.throughput, | ||
| 328 | - CreateExpr(1)); | ||
| 329 | -} | ||
| 330 | - | ||
| 331 | -af::Status GetE2BCount(const BroadcastTilingInfo &tiling, const Expr &vl, Expr &count) { | ||
| 332 | - const size_t last = tiling.rank - 1U; | ||
| 333 | - const Expr last_size = tiling.dst_shape[last]; | ||
| 334 | - GE_ASSERT_TRUE(af::SymbolicUtils::StaticCheckEq(last_size, CreateExpr(0)) != af::TriBool::kTrue, | ||
| 335 | - "Broadcast E2B last-axis size cannot be zero."); | ||
| 336 | - const Expr factor = vl / last_size; | ||
| 337 | - if (tiling.rank == 2U) { | ||
| 338 | - return ascendcapi_v2::CeilDiv(tiling.dst_shape[0], factor, count); | ||
| 339 | - } | ||
| 340 | - if (tiling.rank == 3U) { | ||
| 341 | - if (af::SymbolicUtils::StaticCheckGt(tiling.dst_shape[1] * last_size, vl) == af::TriBool::kTrue) { | ||
| 342 | - Expr inner; | ||
| 343 | - GE_ASSERT_SUCCESS(ascendcapi_v2::CeilDiv(tiling.dst_shape[1], factor, inner)); | ||
| 344 | - count = tiling.dst_shape[0] * inner; | ||
| 345 | - } else { | ||
| 346 | - count = tiling.dst_shape[0]; | ||
| 347 | - } | ||
| 348 | - return af::SUCCESS; | ||
| 349 | - } | ||
| 350 | - Expr inner; | ||
| 351 | - GE_ASSERT_SUCCESS(ascendcapi_v2::CeilDiv(tiling.dst_shape[2], factor, inner)); | ||
| 352 | - count = tiling.dst_shape[0] * tiling.dst_shape[1] * inner; | ||
| 353 | - return af::SUCCESS; | ||
| 354 | -} | ||
| 355 | - | ||
| 356 | -af::Status AddGatherBranch(const NodeDetail &node, const BroadcastTilingInfo &tiling, const Expr &vl, | ||
| 357 | - LastAxisBranch branch, const Expr &outer, VfCostAccumulator &acc) { | ||
| 358 | - const size_t last = tiling.rank - 1U; | ||
| 359 | - Expr gather_count = CreateExpr(1); | ||
| 360 | - if (branch == LastAxisBranch::kGatherTwo) { | ||
| 361 | - const Expr last_size = tiling.dst_shape[last]; | ||
| 362 | - GE_ASSERT_TRUE(af::SymbolicUtils::StaticCheckEq(last_size, CreateExpr(0)) != af::TriBool::kTrue, | ||
| 363 | - "Broadcast last-axis size cannot be zero."); | ||
| 364 | - const Expr factor = vl / last_size; | ||
| 365 | - GE_ASSERT_SUCCESS(ascendcapi_v2::CeilDiv(outer, factor, gather_count)); | ||
| 366 | - } | ||
| 367 | - const Expr lanes = IsB8Dtype(node.input_dtype[0]) ? CreateExpr(2) : CreateExpr(1); | ||
| 368 | - GE_ASSERT_SUCCESS(AddLastGatherIndexPerf(node.input_dtype[0], acc)); | ||
| 369 | - if (branch == LastAxisBranch::kGatherTwo) { | ||
| 370 | - GE_ASSERT_SUCCESS(Add(kDuplicate, GatherRegDtype(node.input_dtype[0]), CreateExpr(1), acc)); | ||
| 371 | - GE_ASSERT_SUCCESS(Add(kMuls, GatherRegDtype(node.input_dtype[0]), gather_count - CreateExpr(1), acc)); | ||
| 372 | - GE_ASSERT_SUCCESS(Add(kAdd, GatherRegDtype(node.input_dtype[0]), (gather_count - CreateExpr(1)) * lanes, acc)); | ||
| 373 | - GE_ASSERT_SUCCESS(Add(kAdds, GatherRegDtype(node.input_dtype[0]), lanes, acc)); | ||
| 374 | - } else { | ||
| 375 | - GE_ASSERT_SUCCESS(Add(kUpdateMask, node.input_dtype[0], CreateExpr(1), acc)); | ||
| 376 | - } | ||
| 377 | - GE_ASSERT_SUCCESS(Add(kLoad, GatherRegDtype(node.input_dtype[0]), lanes, acc)); | ||
| 378 | - // Placeholder: Reg::Gather in last-axis GatherOne/Two, input data dtype, each gather batch. | ||
| 379 | - GE_ASSERT_SUCCESS(VfPerfUtils::AddVfInstructPerf(kPlaceholder, RankTwoGatherDtype(node.input_dtype[0]), | ||
| 380 | - acc.max_latency, acc.throughput, gather_count * lanes)); | ||
| 381 | - if (IsB8Dtype(node.input_dtype[0])) { | ||
| 382 | - GE_ASSERT_SUCCESS(Add(kDeInterleave, node.input_dtype[0], gather_count, acc)); | ||
| 383 | - } | ||
| 384 | - GE_ASSERT_SUCCESS(Add(kStore, node.output_dtype[0], gather_count, acc)); | ||
| 385 | - if (branch == LastAxisBranch::kGatherTwo) { | ||
| 386 | - // Placeholder: Reg::StoreUnAlignPost after BrcLastGatherTwo. | ||
| 387 | - return VfPerfUtils::AddVfInstructPerf(kPlaceholder, node.output_dtype[0], acc.max_latency, acc.throughput, | ||
| 388 | - CreateExpr(1)); | ||
| 389 | - } | ||
| 390 | - return af::SUCCESS; | ||
| 391 | -} | ||
| 392 | - | ||
| 393 | -af::Status AddTailBranch(const NodeDetail &node, LastAxisBranch branch, const Expr &stores, VfCostAccumulator &acc) { | ||
| 394 | - if (branch == LastAxisBranch::kLargerThanVlUnaligned || branch == LastAxisBranch::kLessThanVlUnaligned || | ||
| 395 | - branch == LastAxisBranch::kDynamicLessThanVlUnaligned || | ||
| 396 | - branch == LastAxisBranch::kDynamicLargerThanVlUnaligned) { | ||
| 397 | - GE_ASSERT_SUCCESS(Add(kStore, node.output_dtype[0], stores, acc)); | ||
| 398 | - // Placeholder: Reg::StoreUnAlignPost after the last-axis unaligned stores. | ||
| 399 | - return VfPerfUtils::AddVfInstructPerf(kPlaceholder, node.output_dtype[0], acc.max_latency, acc.throughput, | ||
| 400 | - CreateExpr(1)); | ||
| 401 | - } | ||
| 402 | - return Add(kStore, node.output_dtype[0], stores, acc); | ||
| 403 | -} | ||
| 404 | - | ||
| 405 | -af::Status AddDynamicBranch(const NodeDetail &node, const Expr &outer, const Expr &stores, const Expr &helper_calls, | ||
| 406 | - VfCostAccumulator &acc) { | ||
| 407 | - const std::string helper_dtype = HelperDtype(node.input_dtype[0]); | ||
| 408 | - GE_ASSERT_SUCCESS(Add(kLoad, node.input_dtype[0], outer, acc)); | ||
| 409 | - GE_ASSERT_SUCCESS(Add(kStore, node.output_dtype[0], stores, acc)); | ||
| 410 | - // Placeholder: Reg::StoreUnAlignPost after each dynamic last-axis helper invocation. | ||
| 411 | - return VfPerfUtils::AddVfInstructPerf(kPlaceholder, helper_dtype, acc.max_latency, acc.throughput, helper_calls); | ||
| 412 | -} | ||
| 413 | - | ||
| 414 | -af::Status AddRegularBranch(const NodeDetail &node, const BroadcastTilingInfo &tiling, LastAxisBranch branch, | ||
| 415 | - VfCostAccumulator &acc) { | ||
| 416 | - Expr vl; | ||
| 417 | - Expr half_vl; | ||
| 418 | - Expr block; | ||
| 419 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetBroadcastTilingVectorElements(node.input_dtype[0], vl, half_vl, block)); | ||
| 420 | - (void)half_vl; | ||
| 421 | - (void)block; | ||
| 422 | - const size_t last = tiling.rank - 1U; | ||
| 423 | - const auto &shape = tiling.dst_shape; | ||
| 424 | - Expr outer = ProductFrom(shape, shape.size() - 1U); | ||
| 425 | - if (tiling.loop_num != CreateExpr(0)) { | ||
| 426 | - outer = outer * tiling.loop_num; | ||
| 427 | - } | ||
| 428 | - Expr repeats; | ||
| 429 | - const Expr last_size = tiling.dst_shape[last]; | ||
| 430 | - GE_ASSERT_SUCCESS(ascendcapi_v2::CeilDiv(last_size, vl, repeats)); | ||
| 431 | - const Expr stores = outer * repeats; | ||
| 432 | - if (branch == LastAxisBranch::kDynamicLessThanVlUnaligned || | ||
| 433 | - branch == LastAxisBranch::kDynamicLargerThanVlUnaligned) { | ||
| 434 | - Expr helper_calls = CreateExpr(1); | ||
| 435 | - if (tiling.dst_shape.size() > 4U) { | ||
| 436 | - helper_calls = ProductFrom(tiling.dst_shape, tiling.dst_shape.size() - 4U); | ||
| 437 | - } | ||
| 438 | - return AddDynamicBranch(node, outer, stores, helper_calls, acc); | ||
| 439 | - } | ||
| 440 | - if (branch == LastAxisBranch::kE2B || branch == LastAxisBranch::kE2BLessThanVl || | ||
| 441 | - branch == LastAxisBranch::kE2BLargerThanVl) { | ||
| 442 | - Expr e2b_count; | ||
| 443 | - GE_ASSERT_SUCCESS(GetE2BCount(tiling, vl, e2b_count)); | ||
| 444 | - GE_ASSERT_SUCCESS(Add(kUpdateMask, node.input_dtype[0], e2b_count, acc)); | ||
| 445 | - GE_ASSERT_SUCCESS(Add(kLoad, node.input_dtype[0], e2b_count, acc)); | ||
| 446 | - return Add(kStore, node.output_dtype[0], e2b_count, acc); | ||
| 447 | - } | ||
| 448 | - if (branch == LastAxisBranch::kGatherOne || branch == LastAxisBranch::kGatherTwo) { | ||
| 449 | - return AddGatherBranch(node, tiling, vl, branch, outer, acc); | ||
| 450 | - } | ||
| 451 | - if (branch == LastAxisBranch::kGatherWrapper || branch == LastAxisBranch::kGatherWrapperForFourDim) { | ||
| 452 | - return AddGatherWrapperPerf(node, tiling, branch == LastAxisBranch::kGatherWrapper ? 3U : 4U, acc); | ||
| 453 | - } | ||
| 454 | - if (branch == LastAxisBranch::kLessThanVlAligned) { | ||
| 455 | - GE_ASSERT_SUCCESS(Add(kUpdateMask, node.input_dtype[0], CreateExpr(1), acc)); | ||
| 456 | - } | ||
| 457 | - if (branch == LastAxisBranch::kLargerThanVlAligned) { | ||
| 458 | - const Expr updates = tiling.rank == 4U ? tiling.dst_shape[0] * tiling.dst_shape[2] * repeats : stores; | ||
| 459 | - GE_ASSERT_SUCCESS(Add(kUpdateMask, node.input_dtype[0], updates, acc)); | ||
| 460 | - } | ||
| 461 | - const Expr loads = tiling.rank == 4U ? stores : outer; | ||
| 462 | - GE_ASSERT_SUCCESS(Add(kLoad, node.input_dtype[0], loads, acc)); | ||
| 463 | - return AddTailBranch(node, branch, stores, acc); | ||
| 464 | -} | ||
| 465 | - | ||
| 466 | -af::TriBool Check(const Expr &lhs, CondType condition, const Expr &rhs) { | ||
| 467 | - if (condition == CondType::K_EQ) { | ||
| 468 | - return af::SymbolicUtils::StaticCheckEq(lhs, rhs); | ||
| 469 | - } | ||
| 470 | - if (condition == CondType::K_LT) { | ||
| 471 | - return af::SymbolicUtils::StaticCheckLt(lhs, rhs); | ||
| 472 | - } | ||
| 473 | - if (condition == CondType::K_LE) { | ||
| 474 | - return af::SymbolicUtils::StaticCheckLe(lhs, rhs); | ||
| 475 | - } | ||
| 476 | - return af::SymbolicUtils::StaticCheckGt(lhs, rhs); | ||
| 477 | -} | ||
| 478 | - | ||
| 479 | -Expr LeafCost(const NodeDetail &node, const BroadcastTilingInfo &tiling, LastAxisBranch branch) { | ||
| 480 | - const bool dynamic_wrapper = node.broadcast_node_params.const_rank == -1 || tiling.original_rank > 4U; | ||
| 481 | - if (dynamic_wrapper) { | ||
| 482 | - if (branch == LastAxisBranch::kLargerThanVlUnaligned) { | ||
| 483 | - branch = LastAxisBranch::kDynamicLargerThanVlUnaligned; | ||
| 484 | - } | ||
| 485 | - if (branch == LastAxisBranch::kLessThanVlUnaligned && tiling.rank >= 3U) { | ||
| 486 | - branch = LastAxisBranch::kDynamicLessThanVlUnaligned; | ||
| 487 | - } | ||
| 488 | - } | ||
| 489 | - if (branch == LastAxisBranch::kGatherWrapper || branch == LastAxisBranch::kGatherWrapperForFourDim || | ||
| 490 | - branch == LastAxisBranch::kGatherOne || branch == LastAxisBranch::kGatherTwo) { | ||
| 491 | - Expr vl; | ||
| 492 | - Expr half_vl; | ||
| 493 | - Expr block; | ||
| 494 | - if (ascendcapi_v2::GetBroadcastTilingVectorElements(node.input_dtype[0], vl, half_vl, block) != af::SUCCESS) { | ||
| 495 | - return CreateExpr(0); | ||
| 496 | - } | ||
| 497 | - if (af::SymbolicUtils::StaticCheckLt(tiling.dst_shape.back(), half_vl) == af::TriBool::kFalse) { | ||
| 498 | - return CreateExpr(0); | ||
| 499 | - } | ||
| 500 | - } | ||
| 501 | - VfCostAccumulator acc; | ||
| 502 | - if (AddRegularBranch(node, tiling, branch, acc) != af::SUCCESS) { | ||
| 503 | - return CreateExpr(0); | ||
| 504 | - } | ||
| 505 | - return ascendcapi_v2::GetVfCost(acc); | ||
| 506 | -} | ||
| 507 | - | ||
| 508 | -af::Status BuildDynamicMoreDimLastTree(const NodeDetail &node, const BroadcastTilingInfo &tiling, const Expr &vl, | ||
| 509 | - TernaryOpMap &ternary_ops, Expr &result) { | ||
| 510 | - const Expr last = tiling.dst_shape.back(); | ||
| 511 | - GE_ASSERT_SUCCESS(ascendcapi_v2::BuildBroadcastTernary( | ||
| 512 | - "broadcast_dynamic_last_le_vl", CondType::K_LE, last, vl, | ||
| 513 | - LeafCost(node, tiling, LastAxisBranch::kDynamicLessThanVlUnaligned), | ||
| 514 | - LeafCost(node, tiling, LastAxisBranch::kDynamicLargerThanVlUnaligned), ternary_ops, result)); | ||
| 515 | - return af::SUCCESS; | ||
| 516 | -} | ||
| 517 | - | ||
| 518 | -af::Status Select(const std::string &name, CondType condition, const Expr &lhs, const Expr &rhs, const Expr &true_value, | ||
| 519 | - const Expr &false_value, TernaryOpMap &ternary_ops, Expr &result) { | ||
| 520 | - return ascendcapi_v2::BuildBroadcastTernary(name, condition, lhs, rhs, true_value, false_value, ternary_ops, result); | ||
| 521 | -} | ||
| 522 | - | ||
| 523 | -af::Status BuildTailTree(const NodeDetail &node, const BroadcastTilingInfo &tiling, const Expr &vl, const Expr &half_vl, | ||
| 524 | - const Expr &block, TernaryOpMap &ternary_ops, Expr &result) { | ||
| 525 | - const size_t rank = tiling.rank; | ||
| 526 | - const Expr last = tiling.dst_shape.back(); | ||
| 527 | - const bool b8 = IsB8Dtype(node.input_dtype[0]); | ||
| 528 | - if (rank == 2U) { | ||
| 529 | - Expr larger_aligned; | ||
| 530 | - GE_ASSERT_SUCCESS(Select("broadcast_rank2_block_aligned", CondType::K_EQ, af::sym::Mod(last, block), CreateExpr(0), | ||
| 531 | - LeafCost(node, tiling, LastAxisBranch::kLargerThanVlAligned), | ||
| 532 | - LeafCost(node, tiling, LastAxisBranch::kLargerThanVlUnaligned), ternary_ops, | ||
| 533 | - larger_aligned)); | ||
| 534 | - Expr less_or_larger; | ||
| 535 | - GE_ASSERT_SUCCESS(Select("broadcast_rank2_last_le_vl", CondType::K_LE, last, vl, | ||
| 536 | - LeafCost(node, tiling, LastAxisBranch::kLessThanVlUnaligned), larger_aligned, ternary_ops, | ||
| 537 | - less_or_larger)); | ||
| 538 | - Expr gather; | ||
| 539 | - GE_ASSERT_SUCCESS(Select("broadcast_rank2_total_lt_vl", CondType::K_LT, tiling.dst_size, vl, | ||
| 540 | - LeafCost(node, tiling, LastAxisBranch::kGatherOne), | ||
| 541 | - LeafCost(node, tiling, LastAxisBranch::kGatherTwo), ternary_ops, gather)); | ||
| 542 | - Expr tail; | ||
| 543 | - GE_ASSERT_SUCCESS(Select("broadcast_rank2_last_lt_half_vl", CondType::K_LT, last, half_vl, gather, less_or_larger, | ||
| 544 | - ternary_ops, tail)); | ||
| 545 | - if (!b8) { | ||
| 546 | - GE_ASSERT_SUCCESS(Select("broadcast_rank2_e2b", CondType::K_EQ, last, block, | ||
| 547 | - LeafCost(node, tiling, LastAxisBranch::kE2B), tail, ternary_ops, result)); | ||
| 548 | - } else { | ||
| 549 | - result = tail; | ||
| 550 | - } | ||
| 551 | - return af::SUCCESS; | ||
| 552 | - } | ||
| 553 | - | ||
| 554 | - Expr aligned; | ||
| 555 | - GE_ASSERT_SUCCESS(Select("broadcast_rank_last_le_vl_aligned", CondType::K_EQ, af::sym::Mod(last, block), | ||
| 556 | - CreateExpr(0), LeafCost(node, tiling, LastAxisBranch::kLessThanVlAligned), | ||
| 557 | - LeafCost(node, tiling, LastAxisBranch::kLessThanVlUnaligned), ternary_ops, aligned)); | ||
| 558 | - Expr tail; | ||
| 559 | - GE_ASSERT_SUCCESS(Select("broadcast_rank_last_le_vl", CondType::K_LE, last, vl, aligned, | ||
| 560 | - LeafCost(node, tiling, LastAxisBranch::kLargerThanVlAligned), ternary_ops, tail)); | ||
| 561 | - Expr larger; | ||
| 562 | - GE_ASSERT_SUCCESS(Select("broadcast_rank_last_block_aligned", CondType::K_EQ, af::sym::Mod(last, block), | ||
| 563 | - CreateExpr(0), LeafCost(node, tiling, LastAxisBranch::kLargerThanVlAligned), | ||
| 564 | - LeafCost(node, tiling, LastAxisBranch::kLargerThanVlUnaligned), ternary_ops, larger)); | ||
| 565 | - GE_ASSERT_SUCCESS( | ||
| 566 | - Select("broadcast_rank_last_le_vl_final", CondType::K_LE, last, vl, tail, larger, ternary_ops, result)); | ||
| 567 | - if (!b8 || rank == 3U) { | ||
| 568 | - Expr gather; | ||
| 569 | - const auto gather_branch = rank == 3U ? LastAxisBranch::kGatherWrapper : LastAxisBranch::kGatherWrapperForFourDim; | ||
| 570 | - GE_ASSERT_SUCCESS(Select("broadcast_rank_last_lt_half_vl", CondType::K_LT, last, half_vl, | ||
| 571 | - LeafCost(node, tiling, gather_branch), result, ternary_ops, gather)); | ||
| 572 | - result = gather; | ||
| 573 | - } | ||
| 574 | - if (!b8 && rank == 3U && tiling.src_stride.size() > 1U && | ||
| 575 | - Check(tiling.src_stride[1], CondType::K_EQ, CreateExpr(0)) == af::TriBool::kFalse) { | ||
| 576 | - const Expr inner = tiling.dst_shape[1] * last; | ||
| 577 | - const Expr block_num = vl / block; | ||
| 578 | - Expr e2b_candidate; | ||
| 579 | - GE_ASSERT_SUCCESS(Select("broadcast_rank3_e2b_size", CondType::K_GT, inner, vl, | ||
| 580 | - LeafCost(node, tiling, LastAxisBranch::kE2BLargerThanVl), | ||
| 581 | - LeafCost(node, tiling, LastAxisBranch::kE2BLessThanVl), ternary_ops, e2b_candidate)); | ||
| 582 | - Expr e2b_size; | ||
| 583 | - GE_ASSERT_SUCCESS(Select("broadcast_rank3_e2b_block_num", CondType::K_EQ, | ||
| 584 | - af::sym::Mod(tiling.dst_shape[1], block_num), CreateExpr(0), e2b_candidate, result, | ||
| 585 | - ternary_ops, e2b_size)); | ||
| 586 | - Expr e2b_after_half; | ||
| 587 | - GE_ASSERT_SUCCESS(Select("broadcast_rank3_e2b_half_vl", CondType::K_GT, inner, half_vl, e2b_size, result, | ||
| 588 | - ternary_ops, e2b_after_half)); | ||
| 589 | - const Expr base_result = result; | ||
| 590 | - Expr e2b_result; | ||
| 591 | - GE_ASSERT_SUCCESS(Select("broadcast_rank3_e2b_last", CondType::K_EQ, last, block, e2b_after_half, base_result, | ||
| 592 | - ternary_ops, e2b_result)); | ||
| 593 | - result = e2b_result; | ||
| 594 | - } | ||
| 595 | - if (!b8 && rank == 4U && tiling.src_stride.size() > 2U && | ||
| 596 | - Check(tiling.src_stride[2], CondType::K_EQ, CreateExpr(0)) == af::TriBool::kFalse) { | ||
| 597 | - const Expr block_num = vl / block; | ||
| 598 | - Expr e2b; | ||
| 599 | - GE_ASSERT_SUCCESS(Select("broadcast_rank4_e2b_block_num", CondType::K_EQ, | ||
| 600 | - af::sym::Mod(tiling.dst_shape[2], block_num), CreateExpr(0), | ||
| 601 | - LeafCost(node, tiling, LastAxisBranch::kE2B), result, ternary_ops, e2b)); | ||
| 602 | - const Expr base_result = result; | ||
| 603 | - Expr e2b_result; | ||
| 604 | - GE_ASSERT_SUCCESS( | ||
| 605 | - Select("broadcast_rank4_e2b_last", CondType::K_EQ, last, block, e2b, base_result, ternary_ops, e2b_result)); | ||
| 606 | - result = e2b_result; | ||
| 607 | - } | ||
| 608 | - return af::SUCCESS; | ||
| 609 | -} | ||
| 610 | - | ||
| 611 | -} // namespace | ||
| 612 | - | ||
| 613 | -LastAxisBranch GetLastAxisPerfBranch(const NodeDetail &node_info) { | ||
| 614 | - const auto ¶ms = node_info.broadcast_node_params; | ||
| 615 | - ParamExprInputs src_inputs; | ||
| 616 | - src_inputs.semantic = node_info.input_dims; | ||
| 617 | - src_inputs.size = node_info.input_dims; | ||
| 618 | - src_inputs.actual_size = node_info.repeats; | ||
| 619 | - ParamExprInputs dst_inputs; | ||
| 620 | - dst_inputs.semantic = node_info.output_dims; | ||
| 621 | - dst_inputs.size = node_info.output_dims; | ||
| 622 | - dst_inputs.actual_size = node_info.output_dims; | ||
| 623 | - BroadcastTilingInfo tiling; | ||
| 624 | - if (ascendcapi_v2::BuildBroadcastTiling(params.src_shape, params.dst_shape, node_info.input_dtype[0], src_inputs, | ||
| 625 | - dst_inputs, tiling) != af::SUCCESS) { | ||
| 626 | - return LastAxisBranch::kFallback; | ||
| 627 | - } | ||
| 628 | - return ascendcapi_v2::GetLastAxisBranch(tiling, node_info.input_dtype[0], node_info.broadcast_node_params.const_rank); | ||
| 629 | -} | ||
| 630 | - | ||
| 631 | -af::Status BuildLastAxisPerf(const NodeDetail &node_info, const BroadcastTilingInfo &tiling, PerfOutputInfo &perf) { | ||
| 632 | - const auto branch = | ||
| 633 | - ascendcapi_v2::GetLastAxisBranch(tiling, node_info.input_dtype[0], node_info.broadcast_node_params.const_rank); | ||
| 634 | - GE_ASSERT_TRUE( | ||
| 635 | - branch != LastAxisBranch::kFallback || tiling.rank <= 4U || node_info.broadcast_node_params.const_rank == -1, | ||
| 636 | - "Unsupported last-axis Broadcast branch."); | ||
| 637 | - Expr result; | ||
| 638 | - const bool dynamic_leaf = | ||
| 639 | - branch == LastAxisBranch::kDynamicLessThanVlUnaligned || branch == LastAxisBranch::kDynamicLargerThanVlUnaligned; | ||
| 640 | - if (dynamic_leaf && node_info.broadcast_node_params.const_rank != -1) { | ||
| 641 | - VfCostAccumulator acc; | ||
| 642 | - GE_ASSERT_SUCCESS(AddRegularBranch(node_info, tiling, branch, acc)); | ||
| 643 | - result = ascendcapi_v2::GetVfCost(acc); | ||
| 644 | - } else if (tiling.rank >= 2U && tiling.rank <= 4U && | ||
| 645 | - !(node_info.broadcast_node_params.const_rank == -1 && tiling.original_rank > 4U)) { | ||
| 646 | - Expr vl; | ||
| 647 | - Expr half_vl; | ||
| 648 | - Expr block; | ||
| 649 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetBroadcastTilingVectorElements(node_info.input_dtype[0], vl, half_vl, block)); | ||
| 650 | - GE_ASSERT_SUCCESS(BuildTailTree(node_info, tiling, vl, half_vl, block, perf.ternary_ops, result)); | ||
| 651 | - } else if (branch == LastAxisBranch::kFallback && node_info.broadcast_node_params.const_rank == -1) { | ||
| 652 | - Expr vl; | ||
| 653 | - Expr half_vl; | ||
| 654 | - Expr block; | ||
| 655 | - GE_ASSERT_SUCCESS(ascendcapi_v2::GetBroadcastTilingVectorElements(node_info.input_dtype[0], vl, half_vl, block)); | ||
| 656 | - GE_ASSERT_SUCCESS(BuildDynamicMoreDimLastTree(node_info, tiling, vl, perf.ternary_ops, result)); | ||
| 657 | - } else { | ||
| 658 | - VfCostAccumulator acc; | ||
| 659 | - GE_ASSERT_TRUE(branch != LastAxisBranch::kFallback, "Unsupported last-axis Broadcast branch."); | ||
| 660 | - GE_ASSERT_SUCCESS(AddRegularBranch(node_info, tiling, branch, acc)); | ||
| 661 | - result = ascendcapi_v2::GetVfCost(acc); | ||
| 662 | - } | ||
| 663 | - result.Simplify(); | ||
| 664 | - perf.pipe_res[PipeType::AIV_VEC] = result; | ||
| 665 | - return af::SUCCESS; | ||
| 666 | -} | ||
| 667 | - | ||
| 668 | -} // namespace ascendcperf_v2 | ||
| 669 | -} // namespace att | ||
| @@ -1,45 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 of the License. | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | -namespace att { | ||
| 13 | -namespace ascendcapi_v2 { | ||
| 14 | - | ||
| 15 | -enum class LastAxisBranch { | ||
| 16 | - kE2B, | ||
| 17 | - kE2BLessThanVl, | ||
| 18 | - kE2BLargerThanVl, | ||
| 19 | - kGatherOne, | ||
| 20 | - kGatherTwo, | ||
| 21 | - kGatherWrapper, | ||
| 22 | - kGatherWrapperForFourDim, | ||
| 23 | - kLessThanVlAligned, | ||
| 24 | - kLessThanVlUnaligned, | ||
| 25 | - kLargerThanVlAligned, | ||
| 26 | - kLargerThanVlUnaligned, | ||
| 27 | - kDynamicLessThanVlUnaligned, | ||
| 28 | - kDynamicLargerThanVlUnaligned, | ||
| 29 | - kFallback, | ||
| 30 | -}; | ||
| 31 | - | ||
| 32 | -LastAxisBranch GetLastAxisBranch(const BroadcastTilingInfo &tiling, const std::string &dtype, int32_t const_rank); | ||
| 33 | - | ||
| 34 | -} // namespace ascendcapi_v2 | ||
| 35 | - | ||
| 36 | -namespace ascendcperf_v2 { | ||
| 37 | - | ||
| 38 | -ascendcapi_v2::LastAxisBranch GetLastAxisPerfBranch(const NodeDetail &node_info); | ||
| 39 | -af::Status BuildLastAxisPerf(const NodeDetail &node_info, const ascendcapi_v2::BroadcastTilingInfo &tiling, | ||
| 40 | - PerfOutputInfo &perf); | ||
| 41 | - | ||
| 42 | -} // namespace ascendcperf_v2 | ||
| 43 | -} // namespace att | ||
| 44 | - | ||
| 45 | - | ||
| @@ -1,50 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 of the License. | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | -namespace att { | ||
| 13 | -namespace ascendcapi_v2 { | ||
| 14 | - | ||
| 15 | -enum class NlastAxisBranch { | ||
| 16 | - kGather, | ||
| 17 | - kGatherWrapperForFourDim, | ||
| 18 | - kGatherOne, | ||
| 19 | - kGatherTwo, | ||
| 20 | - kGatherBOne, | ||
| 21 | - kGatherBTwo, | ||
| 22 | - kLessThanVlAligned, | ||
| 23 | - kLessThanVlUnaligned, | ||
| 24 | - kLargerThanVlAlignedWithBlock, | ||
| 25 | - kLargerThanVlAlignedWithVl, | ||
| 26 | - kLargerThanVlUnaligned, | ||
| 27 | - kDynamicGather, | ||
| 28 | - kDynamicLessThanVlAligned, | ||
| 29 | - kDynamicLessThanVlUnaligned, | ||
| 30 | - kDynamicLargerThanVlAlignedWithBlock, | ||
| 31 | - kDynamicLargerThanVlUnaligned, | ||
| 32 | - kB64MoreDimGather, | ||
| 33 | - kFallback, | ||
| 34 | -}; | ||
| 35 | - | ||
| 36 | -NlastAxisBranch GetNlastAxisBranch(const BroadcastTilingInfo &tiling, const std::string &dtype, int32_t const_rank); | ||
| 37 | - | ||
| 38 | -} // namespace ascendcapi_v2 | ||
| 39 | - | ||
| 40 | -namespace ascendcperf_v2 { | ||
| 41 | - | ||
| 42 | -ascendcapi_v2::NlastAxisBranch GetNlastAxisPerfBranch(const NodeDetail &node_info); | ||
| 43 | -bool HasUnknownNlastCondition(const NodeDetail &node_info, const ascendcapi_v2::BroadcastTilingInfo &tiling); | ||
| 44 | -af::Status BuildNlastAxisPerf(const NodeDetail &node_info, const ascendcapi_v2::BroadcastTilingInfo &tiling, | ||
| 45 | - PerfOutputInfo &perf); | ||
| 46 | - | ||
| 47 | -} // namespace ascendcperf_v2 | ||
| 48 | -} // namespace att | ||
| 49 | - | ||
| 50 | - | ||
| @@ -1,278 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under the terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | -namespace att { | ||
| 16 | -namespace ascendcapi_v2 { | ||
| 17 | -namespace { | ||
| 18 | -constexpr uint32_t kBlockBytes = 32U; | ||
| 19 | -constexpr uint32_t kVectorBytes = 256U; | ||
| 20 | - | ||
| 21 | -af::TriBool CheckCondition(CondType type, const Expr &lhs, const Expr &rhs) { | ||
| 22 | - switch (type) { | ||
| 23 | - case CondType::K_EQ: | ||
| 24 | - return af::SymbolicUtils::StaticCheckEq(lhs, rhs); | ||
| 25 | - case CondType::K_LT: | ||
| 26 | - return af::SymbolicUtils::StaticCheckLt(lhs, rhs); | ||
| 27 | - case CondType::K_GT: | ||
| 28 | - return af::SymbolicUtils::StaticCheckGt(lhs, rhs); | ||
| 29 | - case CondType::K_LE: | ||
| 30 | - return af::SymbolicUtils::StaticCheckLe(lhs, rhs); | ||
| 31 | - case CondType::K_GE: | ||
| 32 | - return af::SymbolicUtils::StaticCheckGe(lhs, rhs); | ||
| 33 | - default: | ||
| 34 | - return af::TriBool::kUnknown; | ||
| 35 | - } | ||
| 36 | -} | ||
| 37 | - | ||
| 38 | -void CollapseRank(std::vector<Expr> &src, std::vector<Expr> &dst) { | ||
| 39 | - std::vector<Expr> collapsed_src; | ||
| 40 | - std::vector<Expr> collapsed_dst; | ||
| 41 | - for (size_t i = 0U; i < src.size();) { | ||
| 42 | - const bool broadcast = src[i] == CreateExpr(1) && dst[i] != CreateExpr(1); | ||
| 43 | - Expr src_product = broadcast ? CreateExpr(1) : src[i]; | ||
| 44 | - Expr dst_product = dst[i]; | ||
| 45 | - ++i; | ||
| 46 | - while (i < src.size() && | ||
| 47 | - ((broadcast && src[i] == CreateExpr(1) && dst[i] != CreateExpr(1)) || (!broadcast && src[i] == dst[i]))) { | ||
| 48 | - dst_product = dst_product * dst[i]; | ||
| 49 | - if (!broadcast) { | ||
| 50 | - src_product = src_product * src[i]; | ||
| 51 | - } | ||
| 52 | - ++i; | ||
| 53 | - } | ||
| 54 | - collapsed_src.push_back(src_product); | ||
| 55 | - collapsed_dst.push_back(dst_product); | ||
| 56 | - } | ||
| 57 | - src = std::move(collapsed_src); | ||
| 58 | - dst = std::move(collapsed_dst); | ||
| 59 | -} | ||
| 60 | - | ||
| 61 | -void BuildStrides(const std::vector<Expr> &src, const std::vector<Expr> &dst, std::vector<Expr> &src_stride, | ||
| 62 | - std::vector<Expr> &dst_stride) { | ||
| 63 | - src_stride.assign(src.size(), CreateExpr(0)); | ||
| 64 | - dst_stride.assign(dst.size(), CreateExpr(0)); | ||
| 65 | - Expr src_step = CreateExpr(1); | ||
| 66 | - Expr dst_step = CreateExpr(1); | ||
| 67 | - for (size_t i = src.size(); i > 0U; --i) { | ||
| 68 | - const size_t index = i - 1U; | ||
| 69 | - dst_stride[index] = dst_step; | ||
| 70 | - dst_step = dst_step * dst[index]; | ||
| 71 | - if (src[index] == CreateExpr(1) && dst[index] != CreateExpr(1)) { | ||
| 72 | - continue; | ||
| 73 | - } | ||
| 74 | - src_stride[index] = src_step; | ||
| 75 | - src_step = src_step * src[index]; | ||
| 76 | - } | ||
| 77 | -} | ||
| 78 | - | ||
| 79 | -af::Status ApplyB64Tiling(const std::string &dtype, BroadcastTilingInfo &tiling) { | ||
| 80 | - if (dtype != kUInt64 && dtype != kInt64) { | ||
| 81 | - return af::SUCCESS; | ||
| 82 | - } | ||
| 83 | - if (tiling.src_size == tiling.dst_size) { | ||
| 84 | - return af::SUCCESS; | ||
| 85 | - } | ||
| 86 | - const size_t last = tiling.rank - 1U; | ||
| 87 | - if (tiling.src_shape[last] == CreateExpr(1) && tiling.dst_shape[last] != CreateExpr(1)) { | ||
| 88 | - if (tiling.rank < 9U) { | ||
| 89 | - tiling.src_shape.push_back(kSymTwo); | ||
| 90 | - tiling.dst_shape.push_back(kSymTwo); | ||
| 91 | - ++tiling.rank; | ||
| 92 | - } else { | ||
| 93 | - tiling.loop_num = tiling.dst_shape[0]; | ||
| 94 | - } | ||
| 95 | - } else { | ||
| 96 | - tiling.src_shape[last] = tiling.src_shape[last] * kSymTwo; | ||
| 97 | - tiling.dst_shape[last] = tiling.dst_shape[last] * kSymTwo; | ||
| 98 | - } | ||
| 99 | - tiling.src_size = tiling.src_size * kSymTwo; | ||
| 100 | - tiling.dst_size = tiling.dst_size * kSymTwo; | ||
| 101 | - return af::SUCCESS; | ||
| 102 | -} | ||
| 103 | - | ||
| 104 | -bool HasParamExprInputs(const std::vector<ascir_param::ParamExprLeaf> &leaves, const ParamExprInputs &inputs) { | ||
| 105 | - return leaves.size() == inputs.semantic.size() && leaves.size() == inputs.size.size() && | ||
| 106 | - leaves.size() == inputs.actual_size.size(); | ||
| 107 | -} | ||
| 108 | -} // namespace | ||
| 109 | - | ||
| 110 | -af::Status GetDtypeByteSize(const std::string &dtype, Expr &byte_size) { | ||
| 111 | - const auto iter = kDataTypeSizeMap.find(dtype); | ||
| 112 | - GE_ASSERT_TRUE(iter != kDataTypeSizeMap.end(), "Unsupported broadcast dtype[%s].", dtype.c_str()); | ||
| 113 | - byte_size = iter->second; | ||
| 114 | - return af::SUCCESS; | ||
| 115 | -} | ||
| 116 | - | ||
| 117 | -af::Status GetEffectiveHelperDtype(const std::string &dtype, std::string &helper_dtype) { | ||
| 118 | - Expr byte_size; | ||
| 119 | - GE_ASSERT_SUCCESS(GetDtypeByteSize(dtype, byte_size)); | ||
| 120 | - helper_dtype = (dtype == kUInt64 || dtype == kInt64) ? kUInt32 : dtype; | ||
| 121 | - return af::SUCCESS; | ||
| 122 | -} | ||
| 123 | - | ||
| 124 | -af::Status GetBroadcastVectorElements(const std::string &dtype, Expr &vl, Expr &half_vl, Expr &block_elements) { | ||
| 125 | - Expr byte_size; | ||
| 126 | - GE_ASSERT_SUCCESS(GetDtypeByteSize(dtype, byte_size)); | ||
| 127 | - GE_ASSERT_TRUE(af::SymbolicUtils::StaticCheckEq(byte_size, CreateExpr(0)) != af::TriBool::kTrue, | ||
| 128 | - "Broadcast dtype byte size cannot be zero."); | ||
| 129 | - vl = CreateExpr(kVectorBytes) / byte_size; | ||
| 130 | - GE_ASSERT_TRUE(2U != 0U, "Broadcast half vector divisor cannot be zero."); | ||
| 131 | - half_vl = vl / kSymTwo; | ||
| 132 | - GE_ASSERT_TRUE(af::SymbolicUtils::StaticCheckEq(byte_size, CreateExpr(0)) != af::TriBool::kTrue, | ||
| 133 | - "Broadcast dtype byte size cannot be zero."); | ||
| 134 | - block_elements = CreateExpr(kBlockBytes) / byte_size; | ||
| 135 | - return af::SUCCESS; | ||
| 136 | -} | ||
| 137 | - | ||
| 138 | -af::Status GetBroadcastTilingVectorElements(const std::string &dtype, Expr &vl, Expr &half_vl, Expr &block_elements) { | ||
| 139 | - std::string helper_dtype; | ||
| 140 | - GE_ASSERT_SUCCESS(GetEffectiveHelperDtype(dtype, helper_dtype)); | ||
| 141 | - return GetBroadcastVectorElements(helper_dtype, vl, half_vl, block_elements); | ||
| 142 | -} | ||
| 143 | - | ||
| 144 | -af::Status CeilDiv(const Expr &value, const Expr &divisor, Expr &result) { | ||
| 145 | - GE_ASSERT_TRUE(af::SymbolicUtils::StaticCheckEq(divisor, CreateExpr(0)) != af::TriBool::kTrue, | ||
| 146 | - "Broadcast ceil divisor cannot be zero."); | ||
| 147 | - result = af::sym::Ceiling(value / divisor); | ||
| 148 | - return af::SUCCESS; | ||
| 149 | -} | ||
| 150 | - | ||
| 151 | -Expr ResolveParamExprLeaf(const ascir_param::ParamExprLeaf &leaf, const Expr &semantic_expr, const Expr &size_expr, | ||
| 152 | - const Expr &actual_size_expr) { | ||
| 153 | - if (leaf.role == ascir_param::ParamExprRole::kSize) { | ||
| 154 | - return size_expr; | ||
| 155 | - } | ||
| 156 | - if (leaf.role == ascir_param::ParamExprRole::kActualSize) { | ||
| 157 | - return actual_size_expr; | ||
| 158 | - } | ||
| 159 | - return semantic_expr; | ||
| 160 | -} | ||
| 161 | - | ||
| 162 | -Expr ShapeProduct(const std::vector<Expr> &shape) { | ||
| 163 | - Expr product = CreateExpr(1); | ||
| 164 | - for (const auto &dim : shape) { | ||
| 165 | - product = product * dim; | ||
| 166 | - } | ||
| 167 | - return product; | ||
| 168 | -} | ||
| 169 | - | ||
| 170 | -af::Status BuildBroadcastTiling(const std::vector<ascir_param::ParamExprLeaf> &src, | ||
| 171 | - const std::vector<ascir_param::ParamExprLeaf> &dst, const std::string &dtype, | ||
| 172 | - const ParamExprInputs &src_inputs, const ParamExprInputs &dst_inputs, | ||
| 173 | - BroadcastTilingInfo &tiling) { | ||
| 174 | - GE_ASSERT_TRUE(!src.empty() && src.size() == dst.size(), "Broadcast shapes must have the same non-zero rank."); | ||
| 175 | - GE_ASSERT_TRUE(HasParamExprInputs(src, src_inputs) && HasParamExprInputs(dst, dst_inputs), | ||
| 176 | - "Broadcast parameter expression inputs must match rank."); | ||
| 177 | - tiling.original_src_shape.clear(); | ||
| 178 | - tiling.original_dst_shape.clear(); | ||
| 179 | - for (size_t i = 0U; i < src.size(); ++i) { | ||
| 180 | - tiling.original_src_shape.push_back( | ||
| 181 | - ResolveParamExprLeaf(src[i], src_inputs.semantic[i], src_inputs.size[i], src_inputs.actual_size[i])); | ||
| 182 | - tiling.original_dst_shape.push_back( | ||
| 183 | - ResolveParamExprLeaf(dst[i], dst_inputs.semantic[i], dst_inputs.size[i], dst_inputs.actual_size[i])); | ||
| 184 | - } | ||
| 185 | - tiling.original_rank = src.size(); | ||
| 186 | - tiling.src_shape = tiling.original_src_shape; | ||
| 187 | - tiling.dst_shape = tiling.original_dst_shape; | ||
| 188 | - tiling.src_size = ShapeProduct(tiling.original_src_shape); | ||
| 189 | - tiling.dst_size = ShapeProduct(tiling.original_dst_shape); | ||
| 190 | - if (tiling.src_shape.size() > 4U) { | ||
| 191 | - CollapseRank(tiling.src_shape, tiling.dst_shape); | ||
| 192 | - } | ||
| 193 | - tiling.folded_rank = tiling.dst_shape.size(); | ||
| 194 | - if (tiling.original_rank == 4U && tiling.dst_shape[0] == CreateExpr(1) && | ||
| 195 | - tiling.original_src_shape[0] == CreateExpr(1)) { | ||
| 196 | - tiling.src_shape.erase(tiling.src_shape.begin()); | ||
| 197 | - tiling.dst_shape.erase(tiling.dst_shape.begin()); | ||
| 198 | - } | ||
| 199 | - tiling.rank = tiling.dst_shape.size(); | ||
| 200 | - GE_ASSERT_SUCCESS(ApplyB64Tiling(dtype, tiling)); | ||
| 201 | - const bool loop_src_stride_zero = | ||
| 202 | - tiling.loop_num != CreateExpr(0) && tiling.src_shape[0] == CreateExpr(1) && tiling.dst_shape[0] != CreateExpr(1); | ||
| 203 | - if (tiling.loop_num != CreateExpr(0)) { | ||
| 204 | - tiling.src_shape.erase(tiling.src_shape.begin()); | ||
| 205 | - tiling.dst_shape.erase(tiling.dst_shape.begin()); | ||
| 206 | - tiling.src_shape.push_back(kSymTwo); | ||
| 207 | - tiling.dst_shape.push_back(kSymTwo); | ||
| 208 | - } | ||
| 209 | - BuildStrides(tiling.src_shape, tiling.dst_shape, tiling.src_stride, tiling.dst_stride); | ||
| 210 | - if (tiling.loop_num != CreateExpr(0)) { | ||
| 211 | - tiling.src_stride.push_back(loop_src_stride_zero ? CreateExpr(0) : ShapeProduct(tiling.src_shape)); | ||
| 212 | - } | ||
| 213 | - return af::SUCCESS; | ||
| 214 | -} | ||
| 215 | - | ||
| 216 | -bool IsLastAxisBroadcast(const BroadcastTilingInfo &tiling) { | ||
| 217 | - return tiling.rank > 0U && tiling.src_stride.size() >= tiling.rank && | ||
| 218 | - tiling.src_stride[tiling.rank - 1U] == CreateExpr(0); | ||
| 219 | -} | ||
| 220 | - | ||
| 221 | -af::Status AddBroadcastDataCopyPerf(const std::string &dtype, const Expr &repeat_time, VfCostAccumulator &acc) { | ||
| 222 | - return AddVfInstructPerf(kLoad, dtype, repeat_time, 1U, acc); | ||
| 223 | -} | ||
| 224 | - | ||
| 225 | -af::Status AddVfInstructPerf(const std::string &instruct, const std::string &dtype, const Expr &repeat_time, | ||
| 226 | - uint32_t instruct_count, VfCostAccumulator &acc) { | ||
| 227 | - std::string helper_dtype; | ||
| 228 | - GE_ASSERT_SUCCESS(GetEffectiveHelperDtype(dtype, helper_dtype)); | ||
| 229 | - if (instruct == kLoad) { | ||
| 230 | - acc.load_count += instruct_count; | ||
| 231 | - } | ||
| 232 | - static const PerfParamTableV2 perf_table; | ||
| 233 | - const auto &entries = perf_table.GetVfInstructPerfTable(instruct); | ||
| 234 | - for (uint32_t i = 0U; i < instruct_count; ++i) { | ||
| 235 | - for (const auto &entry : entries) { | ||
| 236 | - if (std::find(entry.support_data_types.begin(), entry.support_data_types.end(), helper_dtype) == | ||
| 237 | - entry.support_data_types.end()) { | ||
| 238 | - continue; | ||
| 239 | - } | ||
| 240 | - acc.max_latency = af::sym::Max(acc.max_latency, CreateExpr(entry.latency)); | ||
| 241 | - acc.throughput = acc.throughput + CreateExpr(entry.throughput) * repeat_time; | ||
| 242 | - break; | ||
| 243 | - } | ||
| 244 | - } | ||
| 245 | - return af::SUCCESS; | ||
| 246 | -} | ||
| 247 | - | ||
| 248 | -Expr GetVfCost(const VfCostAccumulator &acc, bool include_head) { | ||
| 249 | - Expr cost = acc.max_latency + acc.throughput; | ||
| 250 | - if (include_head) { | ||
| 251 | - static const PerfParamTableV2 perf_table; | ||
| 252 | - cost = cost + perf_table.GetVectorFunctionHeadCost(); | ||
| 253 | - } | ||
| 254 | - cost.Simplify(); | ||
| 255 | - return cost; | ||
| 256 | -} | ||
| 257 | - | ||
| 258 | -af::Status BuildBroadcastTernary(const std::string &name, CondType condition, const Expr &lhs, const Expr &rhs, | ||
| 259 | - const Expr &true_value, const Expr &false_value, TernaryOpMap &ternary_ops, | ||
| 260 | - Expr &result) { | ||
| 261 | - const auto check = CheckCondition(condition, lhs, rhs); | ||
| 262 | - if (check == af::TriBool::kTrue) { | ||
| 263 | - result = true_value; | ||
| 264 | - return af::SUCCESS; | ||
| 265 | - } | ||
| 266 | - if (check == af::TriBool::kFalse) { | ||
| 267 | - result = false_value; | ||
| 268 | - return af::SUCCESS; | ||
| 269 | - } | ||
| 270 | - GetPerfVar(name, result, ternary_ops); | ||
| 271 | - TernaryOp ternary(condition, lhs, rhs, true_value, false_value); | ||
| 272 | - ternary.SetVariable(result); | ||
| 273 | - ternary_ops[result] = ternary; | ||
| 274 | - return af::SUCCESS; | ||
| 275 | -} | ||
| 276 | - | ||
| 277 | -} // namespace ascendcapi_v2 | ||
| 278 | -} // namespace att | ||
| @@ -1,74 +0,0 @@ | |||
| 1 | -/** | ||
| 2 | - * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | - * This program is free software, you can redistribute it and/or modify it under the terms of | ||
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | - */ | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | - | ||
| 10 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | -namespace att { | ||
| 20 | -namespace ascendcapi_v2 { | ||
| 21 | - | ||
| 22 | -struct BroadcastTilingInfo { | ||
| 23 | - std::vector<Expr> original_src_shape; | ||
| 24 | - std::vector<Expr> original_dst_shape; | ||
| 25 | - std::vector<Expr> src_shape; | ||
| 26 | - std::vector<Expr> dst_shape; | ||
| 27 | - std::vector<Expr> src_stride; | ||
| 28 | - std::vector<Expr> dst_stride; | ||
| 29 | - Expr src_size{CreateExpr(1)}; | ||
| 30 | - Expr dst_size{CreateExpr(1)}; | ||
| 31 | - Expr loop_num{CreateExpr(0)}; | ||
| 32 | - size_t original_rank{0U}; | ||
| 33 | - size_t folded_rank{0U}; | ||
| 34 | - size_t rank{0U}; | ||
| 35 | -}; | ||
| 36 | - | ||
| 37 | -struct VfCostAccumulator { | ||
| 38 | - Expr max_latency{CreateExpr(0)}; | ||
| 39 | - Expr throughput{CreateExpr(0)}; | ||
| 40 | - uint32_t load_count{0U}; | ||
| 41 | -}; | ||
| 42 | - | ||
| 43 | -struct ParamExprInputs { | ||
| 44 | - std::vector<Expr> semantic; | ||
| 45 | - std::vector<Expr> size; | ||
| 46 | - std::vector<Expr> actual_size; | ||
| 47 | -}; | ||
| 48 | - | ||
| 49 | -af::Status GetDtypeByteSize(const std::string &dtype, Expr &byte_size); | ||
| 50 | -af::Status GetEffectiveHelperDtype(const std::string &dtype, std::string &helper_dtype); | ||
| 51 | -af::Status GetBroadcastVectorElements(const std::string &dtype, Expr &vl, Expr &half_vl, Expr &block_elements); | ||
| 52 | -af::Status GetBroadcastTilingVectorElements(const std::string &dtype, Expr &vl, Expr &half_vl, Expr &block_elements); | ||
| 53 | -af::Status CeilDiv(const Expr &value, const Expr &divisor, Expr &result); | ||
| 54 | -Expr ResolveParamExprLeaf(const ascir_param::ParamExprLeaf &leaf, const Expr &semantic_expr, const Expr &size_expr, | ||
| 55 | - const Expr &actual_size_expr); | ||
| 56 | -Expr ShapeProduct(const std::vector<Expr> &shape); | ||
| 57 | -af::Status BuildBroadcastTiling(const std::vector<ascir_param::ParamExprLeaf> &src, | ||
| 58 | - const std::vector<ascir_param::ParamExprLeaf> &dst, const std::string &dtype, | ||
| 59 | - const ParamExprInputs &src_inputs, const ParamExprInputs &dst_inputs, | ||
| 60 | - BroadcastTilingInfo &tiling); | ||
| 61 | -bool IsLastAxisBroadcast(const BroadcastTilingInfo &tiling); | ||
| 62 | - | ||
| 63 | -af::Status AddVfInstructPerf(const std::string &instruct, const std::string &dtype, const Expr &repeat_time, | ||
| 64 | - uint32_t instruct_count, VfCostAccumulator &acc); | ||
| 65 | -af::Status AddBroadcastDataCopyPerf(const std::string &dtype, const Expr &repeat_time, VfCostAccumulator &acc); | ||
| 66 | -Expr GetVfCost(const VfCostAccumulator &acc, bool include_head = true); | ||
| 67 | -af::Status BuildBroadcastTernary(const std::string &name, CondType condition, const Expr &lhs, const Expr &rhs, | ||
| 68 | - const Expr &true_value, const Expr &false_value, TernaryOpMap &ternary_ops, | ||
| 69 | - Expr &result); | ||
| 70 | - | ||
| 71 | -} // namespace ascendcapi_v2 | ||
| 72 | -} // namespace att | ||
| 73 | - | ||
| 74 | - | ||
| @@ -10,7 +10,6 @@ | |||
| 10 | 10 | ||
| 11 | 11 | ||
| 12 | 12 | ||
| 13 | - | ||
| 14 | 13 | ||
| 15 | 14 | ||
| 16 | 15 | ||
| @@ -557,65 +556,6 @@ af::Status TransposeApi([[maybe_unused]] const std::vector<TensorShapeInfo> &inp | |||
| 557 | GE_ASSERT_SUCCESS(ascendcperf_v2::TransposePerf(node_info, perf_res)); | 556 | GE_ASSERT_SUCCESS(ascendcperf_v2::TransposePerf(node_info, perf_res)); |
| 558 | return af::SUCCESS; | 557 | return af::SUCCESS; |
| 559 | } | 558 | } |
| 560 | - | ||
| 561 | -af::Status ValidateBroadcastBasic(const std::vector<TensorShapeInfo> &input_shapes, | ||
| 562 | - const std::vector<TensorShapeInfo> &output_shapes) { | ||
| 563 | - GE_ASSERT_TRUE(input_shapes.size() == 1U && output_shapes.size() == 1U, | ||
| 564 | - "Broadcast requires exactly one input and one output shape."); | ||
| 565 | - GE_ASSERT_TRUE(!input_shapes[0].data_type.empty() && !output_shapes[0].data_type.empty(), | ||
| 566 | - "Broadcast input/output dtype is missing."); | ||
| 567 | - return af::SUCCESS; | ||
| 568 | -} | ||
| 569 | - | ||
| 570 | -af::Status ValidateBroadcastShapes(const std::vector<TensorShapeInfo> &input_shapes, | ||
| 571 | - const std::vector<TensorShapeInfo> &output_shapes, const NodeInfo &node) { | ||
| 572 | - if (node.broadcast_node_params.is_scalar) { | ||
| 573 | - GE_ASSERT_TRUE(node.broadcast_node_params.duplicate_count.IsValid(), "Broadcast scalar actual size is missing."); | ||
| 574 | - return af::SUCCESS; | ||
| 575 | - } | ||
| 576 | - GE_ASSERT_TRUE(!output_shapes[0].dims.empty(), "Broadcast output shape is missing."); | ||
| 577 | - GE_ASSERT_TRUE(!input_shapes[0].dims.empty()); | ||
| 578 | - GE_ASSERT_TRUE(input_shapes[0].dims.size() == output_shapes[0].dims.size(), | ||
| 579 | - "Broadcast input/output shape ranks are different."); | ||
| 580 | - GE_ASSERT_TRUE(input_shapes[0].repeats.size() == input_shapes[0].dims.size(), | ||
| 581 | - "Broadcast input repeats length is invalid."); | ||
| 582 | - return af::SUCCESS; | ||
| 583 | -} | ||
| 584 | - | ||
| 585 | -af::Status LegacyBroadcastPerf(const std::vector<TensorShapeInfo> &input_shapes, | ||
| 586 | - const std::vector<TensorShapeInfo> &output_shapes, const NodeInfo &node, | ||
| 587 | - PerfOutputInfo &perf_res) { | ||
| 588 | - const auto legacy_perf = GetPerfFunc(kBroadcast); | ||
| 589 | - GE_ASSERT_NOTNULL(legacy_perf, "Legacy Broadcast performance function is unavailable."); | ||
| 590 | - return legacy_perf(input_shapes, output_shapes, node, perf_res); | ||
| 591 | -} | ||
| 592 | - | ||
| 593 | -af::Status BroadcastApiV2([[maybe_unused]] const std::vector<TensorShapeInfo> &input_shapes, | ||
| 594 | - [[maybe_unused]] const std::vector<TensorShapeInfo> &output_shapes, const NodeInfo &node, | ||
| 595 | - PerfOutputInfo &perf_res) { | ||
| 596 | - GE_ASSERT_SUCCESS(ValidateBroadcastBasic(input_shapes, output_shapes)); | ||
| 597 | - if (!node.broadcast_node_params.valid || | ||
| 598 | - (!node.broadcast_node_params.is_scalar && | ||
| 599 | - (node.broadcast_node_params.src_shape.empty() || node.broadcast_node_params.dst_shape.empty()))) { | ||
| 600 | - return LegacyBroadcastPerf(input_shapes, output_shapes, node, perf_res); | ||
| 601 | - } | ||
| 602 | - GE_ASSERT_SUCCESS(ValidateBroadcastShapes(input_shapes, output_shapes, node)); | ||
| 603 | - NodeDetail node_info; | ||
| 604 | - GE_ASSERT_SUCCESS(SetNodeDetail(input_shapes, output_shapes, node_info)); | ||
| 605 | - node_info.broadcast_node_params = node.broadcast_node_params; | ||
| 606 | - if (!input_shapes.empty()) { | ||
| 607 | - node_info.repeats = input_shapes[0].repeats; | ||
| 608 | - } | ||
| 609 | - const auto new_status = ascendcperf_v2::BroadcastPerf(node_info, perf_res); | ||
| 610 | - if (new_status == af::SUCCESS) { | ||
| 611 | - return af::SUCCESS; | ||
| 612 | - } | ||
| 613 | - if (!ascendcperf_v2::IsBroadcastFallback(node_info)) { | ||
| 614 | - return new_status; | ||
| 615 | - } | ||
| 616 | - // Only a classified unsupported implementation branch is delegated to the legacy model. | ||
| 617 | - return LegacyBroadcastPerf(input_shapes, output_shapes, node, perf_res); | ||
| 618 | -} | ||
| 619 | } // namespace ascir_v2 | 559 | } // namespace ascir_v2 |
| 620 | 560 | ||
| 621 | REGISTER_EVAL_FUNC_TAG(kStore, V2, ascir_v2::StoreApiV2); | 561 | REGISTER_EVAL_FUNC_TAG(kStore, V2, ascir_v2::StoreApiV2); |
| @@ -659,7 +599,6 @@ REGISTER_EVAL_FUNC_TAG(kCast, V2, ascir_v2::CastApi); | |||
| 659 | REGISTER_EVAL_FUNC_TAG(kSum, V2, ascir_reduce_v2::SumApi); | 599 | REGISTER_EVAL_FUNC_TAG(kSum, V2, ascir_reduce_v2::SumApi); |
| 660 | REGISTER_EVAL_FUNC_TAG(kRemovePad, V2, ascir_v2::RemovePadApi); | 600 | REGISTER_EVAL_FUNC_TAG(kRemovePad, V2, ascir_v2::RemovePadApi); |
| 661 | REGISTER_EVAL_FUNC_TAG(kWhere, V2, ascir_v2::WhereApi); | 601 | REGISTER_EVAL_FUNC_TAG(kWhere, V2, ascir_v2::WhereApi); |
| 662 | -REGISTER_EVAL_FUNC_TAG(kBroadcast, V2, ascir_v2::BroadcastApiV2); | ||
| 663 | REGISTER_EVAL_FUNC_TAG(kPow, V2, ascir_v2::PowApi); | 602 | REGISTER_EVAL_FUNC_TAG(kPow, V2, ascir_v2::PowApi); |
| 664 | REGISTER_EVAL_FUNC_TAG(kErf, V2, ascir_v2::ErfApi); | 603 | REGISTER_EVAL_FUNC_TAG(kErf, V2, ascir_v2::ErfApi); |
| 665 | REGISTER_EVAL_FUNC_TAG(kTanh, V2, ascir_v2::TanhApi); | 604 | REGISTER_EVAL_FUNC_TAG(kTanh, V2, ascir_v2::TanhApi); |
| @@ -684,8 +623,7 @@ ApiPerfRegister<ApiPerf> indirect_load_api_perf_v2(ApiPerfRegisterV2(kIndirectLo | |||
| 684 | &tiling_schedule_config_table_v2)); | 623 | &tiling_schedule_config_table_v2)); |
| 685 | ApiPerfRegister<ApiPerf> abs_api_perf_v2(ApiPerfRegisterV2(kAbs, kAbs + "V2", nullptr, &perf_param_table_v2, | 624 | ApiPerfRegister<ApiPerf> abs_api_perf_v2(ApiPerfRegisterV2(kAbs, kAbs + "V2", nullptr, &perf_param_table_v2, |
| 686 | &tiling_schedule_config_table_v2)); | 625 | &tiling_schedule_config_table_v2)); |
| 687 | -ApiPerfRegister<ApiPerf> broadcast_api_perf_v2(ApiPerfRegisterV2(kBroadcast, GetPerfFunc(kBroadcast + "V2"), nullptr, | 626 | +ApiPerfRegister<ApiPerf> broadcast_api_perf_v2(ApiPerfRegisterV2(kBroadcast, kBroadcast, nullptr, &perf_param_table_v2, |
| 688 | - &perf_param_table_v2, | ||
| 689 | &tiling_schedule_config_table_v2)); | 627 | &tiling_schedule_config_table_v2)); |
| 690 | ApiPerfRegister<ApiPerf> cast_api_perf_v2(ApiPerfRegisterV2(kCast, kCast + "V2", nullptr, &perf_param_table_v2, | 628 | ApiPerfRegister<ApiPerf> cast_api_perf_v2(ApiPerfRegisterV2(kCast, kCast + "V2", nullptr, &perf_param_table_v2, |
| 691 | &tiling_schedule_config_table_v2)); | 629 | &tiling_schedule_config_table_v2)); |
| @@ -15,7 +15,6 @@ | |||
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | - | ||
| 19 | 18 | ||
| 20 | 19 | ||
| 21 | 20 | ||
| @@ -23,115 +22,15 @@ namespace codegen { | |||
| 23 | using namespace std; | 22 | using namespace std; |
| 24 | using namespace ascgen_utils; | 23 | using namespace ascgen_utils; |
| 25 | 24 | ||
| 26 | -namespace { | ||
| 27 | -constexpr const char *kAscirNodeParams = "AscirNodeParams"; | ||
| 28 | -constexpr const char *kCompactPddingMode = "AscendC::PaddingMode::Compact"; | ||
| 29 | - | ||
| 30 | -af::Status ResetBroadcastNodeParams(const af::AscNodePtr &node, ascir_param::BroadcastNodeParams *&broadcast_params) { | ||
| 31 | - GE_ASSERT_NOTNULL(node); | ||
| 32 | - auto params = ascir_param::GetAscirNodeParams(node); | ||
| 33 | - if (params == nullptr) { | ||
| 34 | - auto op_desc = node->GetOpDesc(); | ||
| 35 | - GE_ASSERT_NOTNULL(op_desc); | ||
| 36 | - params = std::make_shared<ascir_param::AscirNodeParams>(); | ||
| 37 | - GE_ASSERT_TRUE(op_desc->SetExtAttr(kAscirNodeParams, params), "Node:%s SetExtAttr failed", node->GetNamePtr()); | ||
| 38 | - } | ||
| 39 | - auto *specific_params = std::get_if<ascir_param::BroadcastNodeParams>(¶ms->specific_params); | ||
| 40 | - if (specific_params == nullptr) { | ||
| 41 | - params->specific_params = ascir_param::BroadcastNodeParams{}; | ||
| 42 | - specific_params = std::get_if<ascir_param::BroadcastNodeParams>(¶ms->specific_params); | ||
| 43 | - } | ||
| 44 | - GE_ASSERT_NOTNULL(specific_params, "Broadcast specific params is null, node[%s].", node->GetNamePtr()); | ||
| 45 | - params->api_name = node->GetType(); | ||
| 46 | - params->status = ascir_param::ParamBuildStatus::kBuilt; | ||
| 47 | - *specific_params = ascir_param::BroadcastNodeParams{}; | ||
| 48 | - broadcast_params = specific_params; | ||
| 49 | - return af::SUCCESS; | ||
| 50 | -} | ||
| 51 | - | ||
| 52 | -std::vector<ascir_param::ParamExprLeaf> GetBroadcastShape(const TPipe &tpipe, const Tensor &input, const Tensor &output, | ||
| 53 | - bool is_src, bool is_compact_padding) { | ||
| 54 | - std::vector<ascir_param::ParamExprLeaf> shape; | ||
| 55 | - const auto axis_count = input.vectorized_axis.size(); | ||
| 56 | - shape.reserve(axis_count); | ||
| 57 | - for (size_t pos = 0UL; pos < axis_count; ++pos) { | ||
| 58 | - if (is_src && | ||
| 59 | - af::SymbolicUtils::StaticCheckEq(input.axis_size[input.vectorized_axis_pos[pos]], | ||
| 60 | - output.axis_size[output.vectorized_axis_pos[pos]]) != af::TriBool::kTrue) { | ||
| 61 | - shape.push_back({af::Symbol(1U), ascir_param::ParamExprRole::kSemantic}); | ||
| 62 | - continue; | ||
| 63 | - } | ||
| 64 | - const auto axis_id = output.vectorized_axis[pos]; | ||
| 65 | - auto role = tpipe.tiler.GetAxis(axis_id).type != ascir::Axis::Type::kAxisTypeTileInner || | ||
| 66 | - output.vectorized_axis[0] == axis_id | ||
| 67 | - ? ascir_param::ParamExprRole::kActualSize | ||
| 68 | - : ascir_param::ParamExprRole::kSize; | ||
| 69 | - if (pos != axis_count - 1UL) { | ||
| 70 | - shape.push_back({output.axis_size[output.vectorized_axis_pos[pos]], role}); | ||
| 71 | - continue; | ||
| 72 | - } | ||
| 73 | - size_t pre_pos = std::numeric_limits<size_t>::max(); | ||
| 74 | - for (size_t i = 0UL; i < pos; ++i) { | ||
| 75 | - if (output.vectorized_strides[i] != 0) { | ||
| 76 | - pre_pos = i; | ||
| 77 | - } | ||
| 78 | - } | ||
| 79 | - const auto expr = pre_pos == std::numeric_limits<size_t>::max() ? output.axis_size[output.vectorized_axis_pos[pos]] | ||
| 80 | - : output.vectorized_strides[pre_pos]; | ||
| 81 | - if (pre_pos != std::numeric_limits<size_t>::max() && is_compact_padding) { | ||
| 82 | - role = ascir_param::ParamExprRole::kActualSize; | ||
| 83 | - } | ||
| 84 | - shape.push_back({expr, role}); | ||
| 85 | - } | ||
| 86 | - return shape; | ||
| 87 | -} | ||
| 88 | - | ||
| 89 | -af::Status GetBroadcastScalarDuplicateCount(const TPipe &tpipe, const Tensor &output, af::Expression &count) { | ||
| 90 | - const auto parse_tiling_expr = [](std::string value) { | ||
| 91 | - size_t pos = 0UL; | ||
| 92 | - while ((pos = value.find("t->", pos)) != std::string::npos) { | ||
| 93 | - value.erase(pos, 3UL); | ||
| 94 | - } | ||
| 95 | - return af::Expression::Parse(value.c_str()); | ||
| 96 | - }; | ||
| 97 | - count = af::Expression::Parse("0"); | ||
| 98 | - if (output.vectorized_strides.size() != output.vectorized_axis.size()) { | ||
| 99 | - count = af::Expression::Parse("1"); | ||
| 100 | - for (const auto axis_pos : output.vectorized_axis_pos) { | ||
| 101 | - const auto axis_size_str = tpipe.tiler.Size(output.axis_size[axis_pos]); | ||
| 102 | - const auto axis_size = parse_tiling_expr(axis_size_str); | ||
| 103 | - GE_ASSERT_TRUE(axis_size.IsValid(), "Broadcast scalar output axis size is invalid: %s", axis_size_str.c_str()); | ||
| 104 | - count = af::sym::Mul(count, axis_size); | ||
| 105 | - } | ||
| 106 | - return af::SUCCESS; | ||
| 107 | - } | ||
| 108 | - for (size_t i = 0UL; i < output.vectorized_axis.size(); ++i) { | ||
| 109 | - if (output.vectorized_strides[i] == 0) { | ||
| 110 | - continue; | ||
| 111 | - } | ||
| 112 | - const auto axis_pos = output.vectorized_axis_pos[i]; | ||
| 113 | - const auto axis_size_str = tpipe.tiler.Size(output.axis_size[axis_pos]); | ||
| 114 | - const auto axis_size = parse_tiling_expr(axis_size_str); | ||
| 115 | - GE_ASSERT_TRUE(axis_size.IsValid(), "Broadcast scalar output axis size is invalid: %s", axis_size_str.c_str()); | ||
| 116 | - auto term = af::sym::Sub(axis_size, af::Expression::Parse("1")); | ||
| 117 | - if (output.vectorized_strides[i] != 1) { | ||
| 118 | - const auto stride = parse_tiling_expr(tpipe.tiler.Size(output.vectorized_strides[i])); | ||
| 119 | - GE_ASSERT_TRUE(stride.IsValid(), "Broadcast scalar output stride is invalid."); | ||
| 120 | - term = af::sym::Mul(term, stride); | ||
| 121 | - } | ||
| 122 | - count = af::sym::Add(count, term); | ||
| 123 | - } | ||
| 124 | - count = af::sym::Add(count, af::Expression::Parse("1")); | ||
| 125 | - return af::SUCCESS; | ||
| 126 | -} | ||
| 127 | - | ||
| 128 | -} // namespace | ||
| 129 | - | ||
| 130 | static void GenParams(const TPipe &tpipe, const Tensor &input, const Tensor &output, std::stringstream &ss, bool is_src, | 25 | static void GenParams(const TPipe &tpipe, const Tensor &input, const Tensor &output, std::stringstream &ss, bool is_src, |
| 131 | - bool is_compact_padding) { | 26 | + bool has_transpose) { |
| 132 | // 只保证在仅对张量尾轴做32B对齐的场景下有效,若对中间轴做了对齐,则还需要增加处理逻辑 | 27 | // 只保证在仅对张量尾轴做32B对齐的场景下有效,若对中间轴做了对齐,则还需要增加处理逻辑 |
| 133 | auto vectorized_axis_size = input.vectorized_axis.size(); | 28 | auto vectorized_axis_size = input.vectorized_axis.size(); |
| 134 | const char *shape_prefix = is_src ? "src_shape_" : "dst_shape_"; | 29 | const char *shape_prefix = is_src ? "src_shape_" : "dst_shape_"; |
| 30 | + DataCopyParams data_copy_param; | ||
| 31 | + (void)CalculateDmaParams(tpipe, output, output, data_copy_param); | ||
| 32 | + constexpr char kCompactPddingMode[] = "AscendC::PaddingMode::Compact"; | ||
| 33 | + std::string padding_mode = GetPaddingMode(output, data_copy_param, has_transpose); | ||
| 135 | ss << "const uint32_t " << shape_prefix << input.id << "_brc_to_" << output.id << "[" << vectorized_axis_size | 34 | ss << "const uint32_t " << shape_prefix << input.id << "_brc_to_" << output.id << "[" << vectorized_axis_size |
| 136 | << "] = {"; | 35 | << "] = {"; |
| 137 | const char *sep = ""; | 36 | const char *sep = ""; |
| @@ -148,6 +47,7 @@ static void GenParams(const TPipe &tpipe, const Tensor &input, const Tensor &out | |||
| 148 | continue; | 47 | continue; |
| 149 | } | 48 | } |
| 150 | } | 49 | } |
| 50 | + // 非尾轴 | ||
| 151 | if (pos != vectorized_axis_size - 1UL) { | 51 | if (pos != vectorized_axis_size - 1UL) { |
| 152 | GetOneAxisSize(tpipe, output, pos, ss); | 52 | GetOneAxisSize(tpipe, output, pos, ss); |
| 153 | ss << ")"; | 53 | ss << ")"; |
| @@ -168,7 +68,7 @@ static void GenParams(const TPipe &tpipe, const Tensor &input, const Tensor &out | |||
| 168 | ascir::AxisId axis_id = output.vectorized_axis[pos]; | 68 | ascir::AxisId axis_id = output.vectorized_axis[pos]; |
| 169 | auto last_dim_size = output.vectorized_strides[pre_pos]; | 69 | auto last_dim_size = output.vectorized_strides[pre_pos]; |
| 170 | if (tpipe.tiler.GetAxis(axis_id).type != ascir::Axis::Type::kAxisTypeTileInner || | 70 | if (tpipe.tiler.GetAxis(axis_id).type != ascir::Axis::Type::kAxisTypeTileInner || |
| 171 | - output.vectorized_axis[0] == axis_id || is_compact_padding) { | 71 | + output.vectorized_axis[0] == axis_id || padding_mode == kCompactPddingMode) { |
| 172 | ss << tpipe.tiler.ActualSize(last_dim_size); | 72 | ss << tpipe.tiler.ActualSize(last_dim_size); |
| 173 | } else { | 73 | } else { |
| 174 | ss << tpipe.tiler.Size(last_dim_size); | 74 | ss << tpipe.tiler.Size(last_dim_size); |
| @@ -185,19 +85,13 @@ Status BroadcastRegApiCall::Generate(const TPipe &tpipe, const std::vector<ascir | |||
| 185 | const auto &x = inputs[0].get(); | 85 | const auto &x = inputs[0].get(); |
| 186 | const auto &y = outputs[0].get(); | 86 | const auto &y = outputs[0].get(); |
| 187 | (void)RegisterBasicDumpParam(this->api_name_, inputs, outputs); | 87 | (void)RegisterBasicDumpParam(this->api_name_, inputs, outputs); |
| 88 | + | ||
| 188 | if (IsBroadcastConstantTensor(x)) { | 89 | if (IsBroadcastConstantTensor(x)) { |
| 189 | int64_t id = -1; | 90 | int64_t id = -1; |
| 190 | BroadcastScalar(tpipe, current_axis, x, y, id, result, false); | 91 | BroadcastScalar(tpipe, current_axis, x, y, id, result, false); |
| 191 | - ascir_param::BroadcastNodeParams *broadcast_params = nullptr; | ||
| 192 | - if (ResetBroadcastNodeParams(this->node, broadcast_params) == af::SUCCESS && broadcast_params != nullptr) { | ||
| 193 | - broadcast_params->valid = true; | ||
| 194 | - broadcast_params->is_scalar = true; | ||
| 195 | - broadcast_params->const_rank = -1; | ||
| 196 | - GE_ASSERT_SUCCESS(GetBroadcastScalarDuplicateCount(tpipe, y, broadcast_params->duplicate_count)); | ||
| 197 | - broadcast_params->duplicate_count_role = ascir_param::ParamExprRole::kActualSize; | ||
| 198 | - } | ||
| 199 | return af::SUCCESS; | 92 | return af::SUCCESS; |
| 200 | } | 93 | } |
| 94 | + | ||
| 201 | size_t min_vectorized_axis_size = 1UL; | 95 | size_t min_vectorized_axis_size = 1UL; |
| 202 | size_t max_vectorized_axis_size = 9UL; | 96 | size_t max_vectorized_axis_size = 9UL; |
| 203 | if (x.vectorized_axis.size() < min_vectorized_axis_size || x.vectorized_axis.size() > max_vectorized_axis_size) { | 97 | if (x.vectorized_axis.size() < min_vectorized_axis_size || x.vectorized_axis.size() > max_vectorized_axis_size) { |
| @@ -217,16 +111,15 @@ Status BroadcastRegApiCall::Generate(const TPipe &tpipe, const std::vector<ascir | |||
| 217 | x.vectorized_axis.size(), y.vectorized_axis.size()); | 111 | x.vectorized_axis.size(), y.vectorized_axis.size()); |
| 218 | return af::FAILED; | 112 | return af::FAILED; |
| 219 | } | 113 | } |
| 114 | + | ||
| 220 | std::stringstream ss; | 115 | std::stringstream ss; |
| 221 | // 生成参数 const uint32_t *dst_shape; | 116 | // 生成参数 const uint32_t *dst_shape; |
| 222 | std::stringstream params_name; | 117 | std::stringstream params_name; |
| 223 | - DataCopyParams data_copy_param; | 118 | + auto has_transpose = IsGraphHasTransposeNode(this->node); |
| 224 | - (void)CalculateDmaParams(tpipe, y, y, data_copy_param); | 119 | + GenParams(tpipe, x, y, ss, false, has_transpose); |
| 225 | - const auto has_transpose = IsGraphHasTransposeNode(this->node); | ||
| 226 | - const bool is_compact_padding = GetPaddingMode(y, data_copy_param, has_transpose) == kCompactPddingMode; | ||
| 227 | - GenParams(tpipe, x, y, ss, false, is_compact_padding); | ||
| 228 | // 生成参数 const uint32_t *src_shape; | 120 | // 生成参数 const uint32_t *src_shape; |
| 229 | - GenParams(tpipe, x, y, ss, true, is_compact_padding); | 121 | + GenParams(tpipe, x, y, ss, true, has_transpose); |
| 122 | + | ||
| 230 | std::string dtype_name; | 123 | std::string dtype_name; |
| 231 | Tensor::DtypeName(x.dtype, dtype_name); | 124 | Tensor::DtypeName(x.dtype, dtype_name); |
| 232 | ss << this->api_name_ << "<" << dtype_name << "," << x.vectorized_axis.size() << ">("; | 125 | ss << this->api_name_ << "<" << dtype_name << "," << x.vectorized_axis.size() << ">("; |
| @@ -234,16 +127,10 @@ Status BroadcastRegApiCall::Generate(const TPipe &tpipe, const std::vector<ascir | |||
| 234 | ss << y << "[" << tpipe.tiler.TensorVectorizedOffset(current_axis, y) << "], "; | 127 | ss << y << "[" << tpipe.tiler.TensorVectorizedOffset(current_axis, y) << "], "; |
| 235 | // 传入参数 const LocalTensor<T> &src; | 128 | // 传入参数 const LocalTensor<T> &src; |
| 236 | ss << x << "[" << tpipe.tiler.TensorVectorizedOffset(current_axis, x) << "], "; | 129 | ss << x << "[" << tpipe.tiler.TensorVectorizedOffset(current_axis, x) << "], "; |
| 130 | + | ||
| 237 | ss << "dst_shape_" << x.id << "_brc_to_" << y.id << ", src_shape_" << x.id << "_brc_to_" << y.id << ");\n"; | 131 | ss << "dst_shape_" << x.id << "_brc_to_" << y.id << ", src_shape_" << x.id << "_brc_to_" << y.id << ");\n"; |
| 238 | 132 | ||
| 239 | result = ss.str(); | 133 | result = ss.str(); |
| 240 | - ascir_param::BroadcastNodeParams *broadcast_params = nullptr; | ||
| 241 | - if (ResetBroadcastNodeParams(this->node, broadcast_params) == af::SUCCESS && broadcast_params != nullptr) { | ||
| 242 | - broadcast_params->valid = true; | ||
| 243 | - broadcast_params->const_rank = static_cast<int32_t>(x.vectorized_axis.size()); | ||
| 244 | - broadcast_params->dst_shape = GetBroadcastShape(tpipe, x, y, false, is_compact_padding); | ||
| 245 | - broadcast_params->src_shape = GetBroadcastShape(tpipe, x, y, true, is_compact_padding); | ||
| 246 | - } | ||
| 247 | return af::SUCCESS; | 134 | return af::SUCCESS; |
| 248 | } | 135 | } |
| 249 | 136 | ||