[Symbolic SI] Refactor Table of Equivalence (#20627)

This commit is contained in:
Evgenya Nugmanova
2023-10-25 09:54:47 +02:00
committed by GitHub
parent dc4240bc61
commit 7874adb58e
2 changed files with 47 additions and 48 deletions
+20 -21
View File
@@ -8,26 +8,23 @@ using namespace ov;
void TableOfEquivalence::set_as_equal(const Dimension& lhs, const Dimension& rhs) {
const auto &l_label = DimensionTracker::get_label(lhs), r_label = DimensionTracker::get_label(rhs);
bool l_known = dimension_table_of_equivalence.count(l_label) && dimension_table_of_equivalence[l_label],
r_known = dimension_table_of_equivalence.count(r_label) && dimension_table_of_equivalence[r_label];
if (l_known && r_known) {
auto soup_l = dimension_table_of_equivalence[l_label];
soup_l->insert(r_label);
auto soup_r = dimension_table_of_equivalence[r_label];
soup_r->insert(l_label);
soup_l->insert(soup_r->begin(), soup_r->end());
soup_r->insert(soup_l->begin(), soup_l->end());
} else {
auto soup = std::make_shared<std::set<label_t>>();
if (l_known)
soup = dimension_table_of_equivalence[l_label];
else if (r_known)
soup = dimension_table_of_equivalence[r_label];
soup->insert(l_label);
soup->insert(r_label);
dimension_table_of_equivalence[l_label] = soup;
dimension_table_of_equivalence[r_label] = soup;
}
if (l_label == ov::no_label || r_label == ov::no_label)
// TODO after value restriction enabling: non labeled dim propagates restriction (if any) to labeled dim
return;
auto get_soup = [](const label_t& label, EqTable& table) -> EqualitySoup {
if (!table.count(label) || !table.at(label))
table[label] = std::make_shared<std::set<label_t>>(std::set<label_t>{label});
return table.at(label);
};
auto l_soup = get_soup(l_label, dimension_table_of_equivalence);
auto r_soup = get_soup(r_label, dimension_table_of_equivalence);
if (r_soup->size() > l_soup->size()) // we would like to minimize number of iterations in the following for-loop
std::swap(l_soup, r_soup);
l_soup->insert(r_soup->begin(), r_soup->end());
for (const auto& label : *r_soup)
dimension_table_of_equivalence[label] = l_soup;
}
const ValTable& TableOfEquivalence::get_value_equivalence_table() const {
@@ -43,7 +40,9 @@ label_t TableOfEquivalence::get_next_label() {
}
bool TableOfEquivalence::are_equal(const Dimension& lhs, const Dimension& rhs) {
const auto &l_label = DimensionTracker::get_label(lhs), r_label = DimensionTracker::get_label(rhs);
if (!DimensionTracker::has_label(lhs) || !DimensionTracker::has_label(rhs))
return false;
const auto &l_label = DimensionTracker::get_label(lhs), &r_label = DimensionTracker::get_label(rhs);
if (l_label == r_label)
return true;
if (dimension_table_of_equivalence.count(l_label) && dimension_table_of_equivalence[l_label])
+27 -27
View File
@@ -127,38 +127,38 @@ TEST(dimension, dimension_equality) {
DimensionTracker dt(te);
// labeling dimensions
Dimension A, B, C;
dt.set_up_for_tracking(A);
dt.set_up_for_tracking(B);
dt.set_up_for_tracking(C);
PartialShape dimensions = PartialShape::dynamic(5); // A, B, C, D, E
for (auto& dimension : dimensions)
dt.set_up_for_tracking(dimension);
// checking labels are unique
EXPECT_NE(DimensionTracker::get_label(A), no_label);
EXPECT_NE(DimensionTracker::get_label(B), no_label);
EXPECT_NE(DimensionTracker::get_label(C), no_label);
EXPECT_NE(DimensionTracker::get_label(A), DimensionTracker::get_label(B));
EXPECT_NE(DimensionTracker::get_label(B), DimensionTracker::get_label(C));
EXPECT_NE(DimensionTracker::get_label(A), DimensionTracker::get_label(C));
EXPECT_EQ(DimensionTracker::get_label(A), DimensionTracker::get_label(A));
EXPECT_EQ(DimensionTracker::get_label(B), DimensionTracker::get_label(B));
EXPECT_EQ(DimensionTracker::get_label(C), DimensionTracker::get_label(C));
for (const auto& dimension : dimensions)
EXPECT_NE(DimensionTracker::get_label(dimension), no_label);
// setting A == B and B == C
te->set_as_equal(A, B);
te->set_as_equal(C, B);
for (const auto& lhs : dimensions) {
for (const auto& rhs : dimensions) {
if (&lhs == &rhs)
continue;
EXPECT_NE(DimensionTracker::get_label(lhs), DimensionTracker::get_label(rhs));
EXPECT_FALSE(te->are_equal(lhs, rhs));
}
}
// expected to see A == B, B == C and A == C
EXPECT_TRUE(te->are_equal(A, B));
EXPECT_TRUE(te->are_equal(A, C));
EXPECT_TRUE(te->are_equal(B, C));
te->set_as_equal(dimensions[0], dimensions[1]); // A == B
te->set_as_equal(dimensions[3], dimensions[4]); // D == E
te->set_as_equal(dimensions[2], dimensions[3]); // C == D
te->set_as_equal(dimensions[1], dimensions[2]); // B == C
// expected to see A == B == C == D == E
for (const auto& lhs : dimensions)
for (const auto& rhs : dimensions)
EXPECT_TRUE(te->are_equal(lhs, rhs));
// clear up all the tracking info
DimensionTracker::reset_tracking_info(A);
DimensionTracker::reset_tracking_info(B);
DimensionTracker::reset_tracking_info(C);
for (auto& dimension : dimensions)
DimensionTracker::reset_tracking_info(dimension);
// expected to have no label
EXPECT_EQ(DimensionTracker::get_label(A), no_label);
EXPECT_EQ(DimensionTracker::get_label(B), no_label);
EXPECT_EQ(DimensionTracker::get_label(C), no_label);
// checking labels are unique
for (const auto& dimension : dimensions)
EXPECT_EQ(DimensionTracker::get_label(dimension), no_label);
}