Skip to content

Instantly share code, notes, and snippets.

@bhargavkulk
Created June 16, 2026 22:17
Show Gist options
  • Select an option

  • Save bhargavkulk/4708f018634fecaf6498964085fed9c5 to your computer and use it in GitHub Desktop.

Select an option

Save bhargavkulk/4708f018634fecaf6498964085fed9c5 to your computer and use it in GitHub Desktop.
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