Created
June 16, 2026 22:17
-
-
Save bhargavkulk/4708f018634fecaf6498964085fed9c5 to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| import unittest | |
| from collections.abc import Sequence | |
| from dataclasses import dataclass | |
| # TODO: path compression | |
| # TODO: reverse lookup | |
| @dataclass | |
| class Term: | |
| pass | |
| @dataclass | |
| class Atom(Term): | |
| label: str | |
| @dataclass | |
| class App(Term): | |
| label: str | |
| children: list[Term] | |
| @dataclass(frozen=True) | |
| class ENode: | |
| label: str | |
| children: tuple[int, ...] | |
| @dataclass | |
| class UnionFind: | |
| parent: list[int] | |
| def add(self) -> int: | |
| self.parent.append(len(self.parent)) | |
| return self.parent[-1] | |
| def find(self, id: int) -> int: | |
| while id != self.parent[id]: | |
| id = self.parent[id] | |
| assert id == self.parent[id] | |
| return id | |
| def union(self, id1: int, id2: int) -> bool: | |
| id1 = self.find(id1) | |
| id2 = self.find(id2) | |
| if id1 == id2: | |
| return False | |
| self.parent[id2] = id1 | |
| return True | |
| def is_equal(self, id1: int, id2: int) -> bool: | |
| return self.find(id1) == self.find(id2) | |
| @dataclass | |
| class EGraph: | |
| uf: UnionFind | |
| map: dict[ENode, int] | |
| def add(self, enode: ENode) -> int: | |
| enode = self.canonicalize(enode) | |
| if enode not in self.map: | |
| self.map[enode] = self.uf.add() | |
| return self.map[enode] | |
| def union(self, id1: int, id2: int) -> bool: | |
| return self.uf.union(id1, id2) | |
| def rebuild(self): | |
| flag = True | |
| while flag: | |
| old_map = self.map | |
| self.map = dict() | |
| for enode, old_id in old_map.items(): | |
| old_id = self.uf.find(old_id) | |
| enode = self.canonicalize(enode) | |
| new_id = self.map.setdefault(enode, old_id) | |
| flag = self.union(new_id, old_id) | |
| def canonicalize(self, enode: ENode) -> ENode: | |
| return ENode( | |
| enode.label, tuple(self.uf.find(child) for child in enode.children) | |
| ) | |
| def instantiate(self, pattern: Term, subst: dict[str, int]) -> int: | |
| match pattern: | |
| case Atom(label): | |
| return subst[label] | |
| case App(label, children): | |
| ids = tuple(self.instantiate(child, subst) for child in children) | |
| enode = ENode(label, ids) | |
| return self.add(enode) | |
| case _: | |
| raise ValueError("unreachable!") | |
| def ematch(self, pattern: Term, id: int): | |
| return self.ematch_recurse(pattern, id, dict()) | |
| def ematch_recurse( | |
| self, pattern: Term, id: int, subst: dict[str, int] | |
| ) -> list[dict[str, int]]: | |
| match pattern: | |
| case Atom(label): | |
| if label not in subst: | |
| subst = {label: id, **subst} | |
| return [subst] | |
| else: | |
| if subst[label] == id: | |
| return [subst] | |
| else: | |
| return [] | |
| case App(label, children): | |
| res = [] | |
| enodes = [ | |
| enode | |
| for enode in self.enodes_in_eclass(id) | |
| if enode.label == label and len(enode.children) == len(children) | |
| ] | |
| for enode in enodes: | |
| todo = [subst] | |
| for pat_child, node_child in zip(children, enode.children): | |
| new_todo: list[dict[str, int]] = [] | |
| for subst_todo in todo: | |
| substs = self.ematch_recurse( | |
| pat_child, node_child, subst_todo | |
| ) | |
| new_todo.extend(substs) | |
| todo = new_todo | |
| res.extend(todo) | |
| return res | |
| case _: | |
| raise NotImplementedError() | |
| def rewrite(self, rewrites: Sequence[tuple[Term, Term]]): | |
| ids = set(self.map.values()) | |
| matches = [] | |
| for rewrite in rewrites: | |
| for id in ids: | |
| substs = self.ematch(rewrite[0], id) | |
| matches.append((substs, id, rewrite[1])) | |
| for substs, id, rhs in matches: | |
| for subst in substs: | |
| new_id = self.instantiate(rhs, subst) | |
| self.union(id, new_id) | |
| def saturate(self, rewrites: Sequence[tuple[Term, Term]]): | |
| self.rebuild() | |
| while True: | |
| len_parent = len(self.uf.parent) | |
| len_map = len(self.map) | |
| self.rewrite(rewrites) | |
| self.rebuild() | |
| if len_parent == len(self.uf.parent) and len_map == len(self.map): | |
| break | |
| def enodes_in_eclass(self, id_of_class: int) -> list[ENode]: | |
| # everything is canonicalized! | |
| return [enode for enode, id in self.map.items() if id_of_class == id] | |
| def is_equivalent(self, id1: int, id2: int) -> bool: | |
| return self.uf.is_equal(id1, id2) | |
| def mk_egraph() -> EGraph: | |
| return EGraph(UnionFind([]), {}) | |
| # python3 -m unittest main.TestEGraph | |
| class TestEGraph(unittest.TestCase): | |
| def test_union_find(self): | |
| uf = UnionFind([]) | |
| a = uf.add() | |
| b = uf.add() | |
| c = uf.add() | |
| self.assertTrue(a != b) | |
| self.assertTrue(b != c) | |
| self.assertTrue(c != a) | |
| uf.union(a, b) | |
| uf.union(b, c) | |
| self.assertTrue(uf.is_equal(a, c)) | |
| def test_egraph_add(self): | |
| enode = ENode("a", ()) | |
| egraph = mk_egraph() | |
| id1 = egraph.add(enode) | |
| id2 = egraph.add(enode) | |
| self.assertEqual(id1, id2) | |
| def test_congruence_closure(self): | |
| egraph = mk_egraph() | |
| a = egraph.add(ENode("a", ())) | |
| b = egraph.add(ENode("b", ())) | |
| fa = egraph.add(ENode("f", (a,))) | |
| fb = egraph.add(ENode("f", (b,))) | |
| egraph.union(a, b) | |
| egraph.rebuild() | |
| self.assertTrue(egraph.is_equivalent(fa, fb)) | |
| def test_eqsat_simple(self): | |
| egraph = mk_egraph() | |
| a = egraph.add(ENode("a", ())) | |
| b = egraph.add(ENode("b", ())) | |
| a_plus_b = egraph.add(ENode("+", (a, b))) | |
| b_plus_a = egraph.add(ENode("+", (b, a))) | |
| # Rewrites: | |
| # a + b <-> b + a | |
| rewrites = [ | |
| (App("+", [Atom("a"), Atom("b")]), App("+", [Atom("b"), Atom("a")])) | |
| ] | |
| egraph.saturate(rewrites) | |
| self.assertTrue(egraph.is_equivalent(a_plus_b, b_plus_a)) | |
| def test_eqsat_difficult(self): | |
| egraph = mk_egraph() | |
| a = egraph.add(ENode("a", ())) | |
| b = egraph.add(ENode("b", ())) | |
| c = egraph.add(ENode("c", ())) | |
| d = egraph.add(ENode("d", ())) | |
| c_plus_d = egraph.add(ENode("+", (c, d))) | |
| b_plus_c_plus_d = egraph.add(ENode("+", (b, c_plus_d))) | |
| lhs = egraph.add(ENode("+", (a, b_plus_c_plus_d))) | |
| b_plus_a = egraph.add(ENode("+", (b, a))) | |
| c_plus_b_plus_a = egraph.add(ENode("+", (c, b_plus_a))) | |
| rhs = egraph.add(ENode("+", (d, c_plus_b_plus_a))) | |
| # Rewrites: | |
| # a + b <-> b + a | |
| # a + (b + c) <-> (a + b) + c | |
| rewrites = [ | |
| (App("+", [Atom("a"), Atom("b")]), App("+", [Atom("b"), Atom("a")])), | |
| ( | |
| App("+", [Atom("a"), App("+", [Atom("b"), Atom("c")])]), | |
| App("+", [App("+", [Atom("a"), Atom("b")]), Atom("c")]), | |
| ), | |
| ( | |
| App("+", [App("+", [Atom("a"), Atom("b")]), Atom("c")]), | |
| App("+", [Atom("a"), App("+", [Atom("b"), Atom("c")])]), | |
| ), | |
| ] | |
| egraph.saturate(rewrites) | |
| self.assertTrue(egraph.is_equivalent(lhs, rhs)) | |
| def test_eqsat_very_difficult(self): | |
| egraph = mk_egraph() | |
| a = egraph.add(ENode("a", ())) | |
| b = egraph.add(ENode("b", ())) | |
| c = egraph.add(ENode("c", ())) | |
| d = egraph.add(ENode("d", ())) | |
| e = egraph.add(ENode("e", ())) | |
| f = egraph.add(ENode("f", ())) | |
| g = egraph.add(ENode("g", ())) | |
| # a + (b + (c + (d + (e + (f + g))))) = g + (f + (e + (d + (c + (b + a))))) | |
| f_plus_g = egraph.add(ENode("+", (f, g))) | |
| e_plus_f_plus_g = egraph.add(ENode("+", (e, f_plus_g))) | |
| d_plus_e_plus_f_plus_g = egraph.add(ENode("+", (d, e_plus_f_plus_g))) | |
| c_plus_d_plus_e_plus_f_plus_g = egraph.add( | |
| ENode("+", (c, d_plus_e_plus_f_plus_g)) | |
| ) | |
| b_plus_c_plus_d_plus_e_plus_f_plus_g = egraph.add( | |
| ENode("+", (b, c_plus_d_plus_e_plus_f_plus_g)) | |
| ) | |
| lhs = egraph.add(ENode("+", (a, b_plus_c_plus_d_plus_e_plus_f_plus_g))) | |
| b_plus_a = egraph.add(ENode("+", (b, a))) | |
| c_plus_b_plus_a = egraph.add(ENode("+", (c, b_plus_a))) | |
| d_plus_c_plus_b_plus_a = egraph.add(ENode("+", (d, c_plus_b_plus_a))) | |
| e_plus_d_plus_c_plus_b_plus_a = egraph.add( | |
| ENode("+", (e, d_plus_c_plus_b_plus_a)) | |
| ) | |
| f_plus_e_plus_d_plus_c_plus_b_plus_a = egraph.add( | |
| ENode("+", (f, e_plus_d_plus_c_plus_b_plus_a)) | |
| ) | |
| rhs = egraph.add(ENode("+", (g, f_plus_e_plus_d_plus_c_plus_b_plus_a))) | |
| # Rewrites: | |
| # a + b <-> b + a | |
| # a + (b + c) <-> (a + b) + c | |
| rewrites = [ | |
| (App("+", [Atom("a"), Atom("b")]), App("+", [Atom("b"), Atom("a")])), | |
| ( | |
| App("+", [Atom("a"), App("+", [Atom("b"), Atom("c")])]), | |
| App("+", [App("+", [Atom("a"), Atom("b")]), Atom("c")]), | |
| ), | |
| ( | |
| App("+", [App("+", [Atom("a"), Atom("b")]), Atom("c")]), | |
| App("+", [Atom("a"), App("+", [Atom("b"), Atom("c")])]), | |
| ), | |
| ] | |
| egraph.saturate(rewrites) | |
| self.assertTrue(egraph.is_equivalent(lhs, rhs)) |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment