已合并
fix: 修复 AutoFuse unit repeat 轴映射 #4551
fix: 修复 AutoFuse unit repeat 轴映射 #4551
已合并
ling-DT创建于 8 天前
2 个文件变更+39-3
@@ -145,10 +145,30 @@ Status BuildUnitRepeatAxisIndex(const std::vector<Expression> &node_repeats,
145 std::vector<int32_t> node_to_base(node_repeats.size(), -1);145 std::vector<int32_t> node_to_base(node_repeats.size(), -1);
146 std::vector<int32_t> base_to_node(base_repeats.size(), -1);146 std::vector<int32_t> base_to_node(base_repeats.size(), -1);
147 147 
148+ const auto same_index_can_map = [&node_repeats, &base_repeats](const size_t index) {
149+ return node_repeats[index] == base_repeats[index] || node_repeats[index] == 1 || base_repeats[index] == 1;
150+ };
151+ const auto reserve_same_index = [&node_to_base, &base_to_node](const size_t index) {
152+ node_to_base[index] = static_cast<int32_t>(index);
153+ base_to_node[index] = static_cast<int32_t>(index);
154+ };
155+ const auto is_unmapped = [&node_to_base, &base_to_node](const int32_t node_index, const int32_t base_index) {
156+ return node_to_base[node_index] == -1 && base_to_node[base_index] == -1;
157+ };
158+ 
159+ for (size_t index = 0U; index < std::min(node_repeats.size(), base_repeats.size()); ++index) {
160+ if (same_index_can_map(index)) {
161+ reserve_same_index(index);
162+ }
163+ }
164+ 
148 for (int32_t node_index = static_cast<int32_t>(node_repeats.size()) - 1; node_index >= 0; --node_index) {165 for (int32_t node_index = static_cast<int32_t>(node_repeats.size()) - 1; node_index >= 0; --node_index) {
166+ if (node_to_base[node_index] != -1) {
167+ continue;
168+ }
149 int32_t matched_base_index = -1;169 int32_t matched_base_index = -1;
150 for (int32_t base_index = static_cast<int32_t>(base_repeats.size()) - 1; base_index >= 0; --base_index) {170 for (int32_t base_index = static_cast<int32_t>(base_repeats.size()) - 1; base_index >= 0; --base_index) {
151- if (base_to_node[base_index] != -1) {171+ if (!is_unmapped(node_index, base_index)) {
152 continue;172 continue;
153 }173 }
154 if (node_repeats[node_index] == base_repeats[base_index]) {174 if (node_repeats[node_index] == base_repeats[base_index]) {
@@ -159,7 +179,7 @@ Status BuildUnitRepeatAxisIndex(const std::vector<Expression> &node_repeats,
159 179 
160 if (matched_base_index == -1) {180 if (matched_base_index == -1) {
161 for (int32_t base_index = static_cast<int32_t>(base_repeats.size()) - 1; base_index >= 0; --base_index) {181 for (int32_t base_index = static_cast<int32_t>(base_repeats.size()) - 1; base_index >= 0; --base_index) {
162- if (base_to_node[base_index] != -1) {182+ if (!is_unmapped(node_index, base_index)) {
163 continue;183 continue;
164 }184 }
165 if ((node_repeats[node_index] == 1) || (base_repeats[base_index] == 1)) {185 if ((node_repeats[node_index] == 1) || (base_repeats[base_index] == 1)) {
@@ -574,10 +574,26 @@ TEST_F(AscGraphAxisMappingTest2, AscGraphAxisMapping_CanAxisMapAllowUnitRepeat_R
574 AscGraphAxisMapping axis_mapping;574 AscGraphAxisMapping axis_mapping;
575 EXPECT_TRUE(axis_mapping.CanAxisMapAllowUnitRepeat(node1_axis, node1_repeats, node2_axis, node2_repeats, node1_map,575 EXPECT_TRUE(axis_mapping.CanAxisMapAllowUnitRepeat(node1_axis, node1_repeats, node2_axis, node2_repeats, node1_map,
576 node2_map, temp_node1_map, temp_node2_map));576 node2_map, temp_node1_map, temp_node2_map));
577- EXPECT_EQ(temp_node1_map, (AxisPairSet{{0, 2}, {1, 4}}));577+ EXPECT_EQ(temp_node1_map, (AxisPairSet{{0, 2}, {1, 3}}));
578 EXPECT_EQ(temp_node2_map, (AxisPairSet{{2, 2}, {3, 3}, {4, 4}}));578 EXPECT_EQ(temp_node2_map, (AxisPairSet{{2, 2}, {3, 3}, {4, 4}}));
579}579}
580 580 
581+TEST_F(AscGraphAxisMappingTest2, AscGraphAxisMapping_CanAxisMapAllowUnitRepeat_PreserveSameIndex) {
582+ std::vector<int64_t> node1_axis{0, 1, 2};
583+ std::vector<Expression> node1_repeats{Symbol(256), Symbol(1), Symbol(50)};
584+ std::vector<int64_t> node2_axis{0, 1, 2};
585+ std::vector<Expression> node2_repeats{Symbol(256), Symbol(1), Symbol(1)};
586+ AxisPairSet node1_map;
587+ AxisPairSet node2_map;
588+ AxisPairSet temp_node1_map;
589+ AxisPairSet temp_node2_map;
590+ 
591+ AscGraphAxisMapping axis_mapping;
592+ EXPECT_TRUE(axis_mapping.CanAxisMapAllowUnitRepeat(node1_axis, node1_repeats, node2_axis, node2_repeats, node1_map,
593+ node2_map, temp_node1_map, temp_node2_map));
594+ EXPECT_EQ(temp_node2_map, (AxisPairSet{{0, 0}, {1, 1}, {2, 2}}));
595+}
596+ 
581TEST_F(AscGraphAxisMappingTest2, AscGraphAxisMapping_CanAxisMapAllowUnitRepeat_ReverseMappingFailed) {597TEST_F(AscGraphAxisMappingTest2, AscGraphAxisMapping_CanAxisMapAllowUnitRepeat_ReverseMappingFailed) {
582 std::vector<int64_t> node1_axis{0, 1};598 std::vector<int64_t> node1_axis{0, 1};
583 std::vector<Expression> node1_repeats{Symbol(2), Symbol(4)};599 std::vector<Expression> node1_repeats{Symbol(2), Symbol(4)};