已合并
性能建模回退 #2052
已合并
高煜博创建于 9月11日
共 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 &params, NodeInfo &
38 return af::SUCCESS;38 return af::SUCCESS;
39}39}
40 40 
41-af::Status FillBroadcastParams(const ascir_param::AscirNodeParams &params, 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- 
48af::Status FillCompareParams(const ascir_param::AscirNodeParams &params, NodeInfo &node_info) {41af::Status FillCompareParams(const ascir_param::AscirNodeParams &params, 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- 
94struct CompareNodeParams {84struct 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 
132using AnySpecificParams =122using 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 
136struct AscirNodeParams {126struct AscirNodeParams {
137 // 扩展属性载荷版本,用于后续兼容。127 // 扩展属性载荷版本,用于后续兼容。
@@ -25,7 +25,6 @@ namespace {
25constexpr const char *kAscirNodeParams = "AscirNodeParams";25constexpr const char *kAscirNodeParams = "AscirNodeParams";
26constexpr const char *kVectorFunc = "VectorFunc";26constexpr const char *kVectorFunc = "VectorFunc";
27constexpr const char *kCast = "Cast";27constexpr const char *kCast = "Cast";
28-constexpr const char *kBroadcast = "Broadcast";
29 28 
30bool IsCompareParamSupported(const std::string &api_name) {29bool 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- 
143af::Status RegisterCompareAscirNodeParams(const af::AscNodePtr &node) {126af::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#define private public17#define private public
18#include "expr_gen/generate_tiling_expr.h"18#include "expr_gen/generate_tiling_expr.h"
19#include "parser/ascend_graph_parser.h"19#include "parser/ascend_graph_parser.h"
20-#include "parser/specific_params_builder.h"
21#undef private20#undef private
22#include "tests/ut/att/utils/graph_construct_utils.h"21#include "tests/ut/att/utils/graph_construct_utils.h"
23 22 
@@ -94,38 +93,6 @@ Status BuildReduceAscendGraphND(AscGraph &graph) {
94} // namespace af93} // namespace af
95namespace att {94namespace att {
96namespace {95namespace {
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-}
129ascir::FusedScheduledResult BuildGatherReduceScheduleResult(const af::AscGraph &gather_graph,96ascir::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- 
194TEST(AscirNodeParamsTest, EnrichReduceParamsForArSingleReduce) {104TEST(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;
@@ -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-#include <string>
8-#include <vector>
9- 
10-#include "gtest/gtest.h"
11-#include "ascir_node_param/ascir_node_param.h"
12-#include "att/api_perf_register/perf_param_v2.h"
13-#include "base/att_const_values.h"
14-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_api_perf_v2.h"
15-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_last_axis_perf_v2.h"
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-#include <string>
8-#include <vector>
9- 
10-#include "gtest/gtest.h"
11-#include "att/api_perf_register/perf_param_v2.h"
12-#include "ascir_node_param/ascir_node_param.h"
13-#include "base/att_const_values.h"
14-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_api_perf_v2.h"
15-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_last_axis_perf_v2.h"
16-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_nlast_axis_perf_v2.h"
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-#include <cstdint>
12-#include <string>
13-#include <tuple>
14-#include <utility>
15-#include <vector>
16- 
17-#include "gtest/gtest.h"
18-#include "ascir_node_param/ascir_node_param.h"
19-#include "common/checker.h"
20-#include "base/att_const_values.h"
21-#include "api_perf_register/api_perf_factory.h"
22-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_api_perf_v2.h"
23-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_perf_utils_v2.h"
24-#include "../../../../../ut/att/testcase/gen_model_info/api_perf_register/runtime_stub.h"
25-#include "tests/depends/slog/src/slog_stub.h"
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#include "node_utils_ex.h"12#include "node_utils_ex.h"
13#include "graph_utils.h"13#include "graph_utils.h"
14#include "ascendc_ir.h"14#include "ascendc_ir.h"
15-#include "ascir_node_param/ascir_node_param.h"
16#include "ascir_ops.h"15#include "ascir_ops.h"
17#include "ascir_ops_utils.h"16#include "ascir_ops_utils.h"
18#include "utils/api_call_factory.h"17#include "utils/api_call_factory.h"
@@ -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-#include "broadcast_api_perf_v2.h"
8- 
9-#include "base/att_const_values.h"
10-#include "common/checker.h"
11-#include "broadcast_perf_utils_v2.h"
12-#include "broadcast_last_axis_perf_v2.h"
13-#include "broadcast_nlast_axis_perf_v2.h"
14-#include "api_perf_register/ascendc_api_perf.h"
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 &params = 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 &params = 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 &params = 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-#ifndef AUTOFUSE_ASCENDC_BROADCAST_API_PERF_V2_H_
8-#define AUTOFUSE_ASCENDC_BROADCAST_API_PERF_V2_H_
9- 
10-#include "api_perf_register/api_perf.h"
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-#endif // AUTOFUSE_ASCENDC_BROADCAST_API_PERF_V2_H_
@@ -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-#include "broadcast_last_axis_perf_v2.h"
8- 
9-#include "base/att_const_values.h"
10-#include "common/checker.h"
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 &params = 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-#ifndef AUTOFUSE_ASCENDC_BROADCAST_LAST_AXIS_PERF_V2_H_
8-#define AUTOFUSE_ASCENDC_BROADCAST_LAST_AXIS_PERF_V2_H_
9- 
10-#include "broadcast_perf_utils_v2.h"
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-#endif // AUTOFUSE_ASCENDC_BROADCAST_LAST_AXIS_PERF_V2_H_
@@ -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-#ifndef AUTOFUSE_ASCENDC_BROADCAST_NLAST_AXIS_PERF_V2_H_
8-#define AUTOFUSE_ASCENDC_BROADCAST_NLAST_AXIS_PERF_V2_H_
9- 
10-#include "broadcast_perf_utils_v2.h"
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-#endif // AUTOFUSE_ASCENDC_BROADCAST_NLAST_AXIS_PERF_V2_H_
@@ -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-#include "broadcast_perf_utils_v2.h"
8- 
9-#include "base/att_const_values.h"
10-#include "common/checker.h"
11-#include "../perf_param_v2.h"
12- 
13-#include <algorithm>
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-#ifndef AUTOFUSE_ASCENDC_BROADCAST_PERF_UTILS_V2_H_
8-#define AUTOFUSE_ASCENDC_BROADCAST_PERF_UTILS_V2_H_
9- 
10-#include <cstddef>
11-#include <cstdint>
12-#include <string>
13-#include <vector>
14- 
15-#include "api_perf_register/api_perf.h"
16-#include "ascir_node_param/ascir_node_param.h"
17-#include "gen_model_info/api_perf_register/utils/vf_perf_utils.h"
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-#endif // AUTOFUSE_ASCENDC_BROADCAST_PERF_UTILS_V2_H_
@@ -10,7 +10,6 @@
10#include "perf_param_v2.h"10#include "perf_param_v2.h"
11#include "nddma_model.h"11#include "nddma_model.h"
12#include "v35/att/api_perf_register/ascir_reduce_api_perf_v2.h"12#include "v35/att/api_perf_register/ascir_reduce_api_perf_v2.h"
13-#include "v35/att/api_perf_register/ascendc_api_perf/broadcast_api_perf_v2.h"
14#include "v35/att/api_perf_register/ascendc_regbase_perf.h"13#include "v35/att/api_perf_register/ascendc_regbase_perf.h"
15#include "api_perf_register/api_perf_factory.h"14#include "api_perf_register/api_perf_factory.h"
16#include "api_perf_register/ascendc_api_perf.h"15#include "api_perf_register/ascendc_api_perf.h"
@@ -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_v2559} // namespace ascir_v2
620 560 
621REGISTER_EVAL_FUNC_TAG(kStore, V2, ascir_v2::StoreApiV2);561REGISTER_EVAL_FUNC_TAG(kStore, V2, ascir_v2::StoreApiV2);
@@ -659,7 +599,6 @@ REGISTER_EVAL_FUNC_TAG(kCast, V2, ascir_v2::CastApi);
659REGISTER_EVAL_FUNC_TAG(kSum, V2, ascir_reduce_v2::SumApi);599REGISTER_EVAL_FUNC_TAG(kSum, V2, ascir_reduce_v2::SumApi);
660REGISTER_EVAL_FUNC_TAG(kRemovePad, V2, ascir_v2::RemovePadApi);600REGISTER_EVAL_FUNC_TAG(kRemovePad, V2, ascir_v2::RemovePadApi);
661REGISTER_EVAL_FUNC_TAG(kWhere, V2, ascir_v2::WhereApi);601REGISTER_EVAL_FUNC_TAG(kWhere, V2, ascir_v2::WhereApi);
662-REGISTER_EVAL_FUNC_TAG(kBroadcast, V2, ascir_v2::BroadcastApiV2);
663REGISTER_EVAL_FUNC_TAG(kPow, V2, ascir_v2::PowApi);602REGISTER_EVAL_FUNC_TAG(kPow, V2, ascir_v2::PowApi);
664REGISTER_EVAL_FUNC_TAG(kErf, V2, ascir_v2::ErfApi);603REGISTER_EVAL_FUNC_TAG(kErf, V2, ascir_v2::ErfApi);
665REGISTER_EVAL_FUNC_TAG(kTanh, V2, ascir_v2::TanhApi);604REGISTER_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));
685ApiPerfRegister<ApiPerf> abs_api_perf_v2(ApiPerfRegisterV2(kAbs, kAbs + "V2", nullptr, &perf_param_table_v2,624ApiPerfRegister<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));
690ApiPerfRegister<ApiPerf> cast_api_perf_v2(ApiPerfRegisterV2(kCast, kCast + "V2", nullptr, &perf_param_table_v2,628ApiPerfRegister<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#include "api_call/utils/api_call_utils.h"15#include "api_call/utils/api_call_utils.h"
16#include "api_call/utils/api_call_factory.h"16#include "api_call/utils/api_call_factory.h"
17#include "api_call/broadcast/broadcast_api_call.h"17#include "api_call/broadcast/broadcast_api_call.h"
18-#include "ascir_node_param/ascir_node_param.h"
19#include "codegen/expression_convert_struct.h"18#include "codegen/expression_convert_struct.h"
20#include "reg_api_call_utils.h"19#include "reg_api_call_utils.h"
21 20 
@@ -23,115 +22,15 @@ namespace codegen {
23using namespace std;22using namespace std;
24using namespace ascgen_utils;23using 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>(&params->specific_params);
40- if (specific_params == nullptr) {
41- params->specific_params = ascir_param::BroadcastNodeParams{};
42- specific_params = std::get_if<ascir_param::BroadcastNodeParams>(&params->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- 
130static void GenParams(const TPipe &tpipe, const Tensor &input, const Tensor &output, std::stringstream &ss, bool is_src,25static 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_size34 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