[Symbolic SI] Refactor Table of Equivalence (#20627)
This commit is contained in:
@@ -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])
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user