已合并
fix: 修复 AutoFuse unit repeat 轴映射 #4551
ling-DT创建于 8 天前
fix: 修复 AutoFuse unit repeat 轴映射 #4551
已合并
共 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 | + | ||
| 581 | TEST_F(AscGraphAxisMappingTest2, AscGraphAxisMapping_CanAxisMapAllowUnitRepeat_ReverseMappingFailed) { | 597 | TEST_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)}; |