From 463f2bc8209595f60fabd3d3695d71223d4fcbe9 Mon Sep 17 00:00:00 2001 From: lubarsky Date: Wed, 30 Sep 2026 02:59:40 +0300 Subject: [PATCH 1/2] Fix quadratic let hoisting when converting deep terms `typed_expr_to_egg(expr, expr_to_let=True)` searches the whole term for shared subexpressions and hoists them into let bindings. Since 14.0.0, `_expr_to_egg` converts every argument with `expr_to_let=True` as well, so the search reran on every subterm and converting a term became quadratic in its depth. A search at the top of a term already hoists every shared subexpression below it, so the arguments now skip it (`lets_hoisted=True`) and only use the let cache. Let bodies are still searched as before, so the emitted program, including synthetic let names and order, is unchanged. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BnJo4ywpA3FufMmo79Av6f --- python/egglog/egraph_state.py | 35 ++++++++++++++++++++------ python/tests/test_high_level.py | 44 +++++++++++++++++++++++++++++++++ 2 files changed, 71 insertions(+), 8 deletions(-) diff --git a/python/egglog/egraph_state.py b/python/egglog/egraph_state.py index ec8312ea..c795b924 100644 --- a/python/egglog/egraph_state.py +++ b/python/egglog/egraph_state.py @@ -1305,16 +1305,25 @@ def typed_expr_to_egg( self, typed_expr_decl: TypedExprDecl, expr_to_let: bool = True, + *, + lets_hoisted: bool = False, ) -> bindings._Expr: + """ + Convert a typed expression to an egg expression. + + If `lets_hoisted` is set, an enclosing call has already hoisted this expression's shared + subexpressions into let bindings, so they are not searched for again. Searching once per + subterm would make converting a term quadratic in its depth. + """ # transform all expressions with multiple parents into a let binding, so that less expressions # are sent to egglog. Only for performance reasons. - if expr_to_let: + if expr_to_let and not lets_hoisted: have_multiple_parents = _exprs_multiple_parents(typed_expr_decl) for expr in reversed(have_multiple_parents): self._transform_let(expr) self.type_ref_to_egg(typed_expr_decl.tp) - return self._expr_to_egg(typed_expr_decl.expr, expr_to_let=expr_to_let) + return self._expr_to_egg(typed_expr_decl.expr, expr_to_let=expr_to_let, lets_hoisted=expr_to_let) def _transform_let(self, typed_expr: TypedExprDecl) -> TypedExprDecl | None: """ @@ -1342,17 +1351,27 @@ def _transform_let(self, typed_expr: TypedExprDecl) -> TypedExprDecl | None: return None @overload - def _expr_to_egg(self, expr_decl: CallDecl, *, expr_to_let: bool = ...) -> bindings.Call: ... + def _expr_to_egg( + self, expr_decl: CallDecl, *, expr_to_let: bool = ..., lets_hoisted: bool = ... + ) -> bindings.Call: ... @overload - def _expr_to_egg(self, expr_decl: UnboundVarDecl | LetRefDecl, *, expr_to_let: bool = ...) -> bindings.Var: ... + def _expr_to_egg( + self, expr_decl: UnboundVarDecl | LetRefDecl, *, expr_to_let: bool = ..., lets_hoisted: bool = ... + ) -> bindings.Var: ... @overload - def _expr_to_egg(self, expr_decl: ExprDecl, *, expr_to_let: bool = ...) -> bindings._Expr: ... + def _expr_to_egg( + self, expr_decl: ExprDecl, *, expr_to_let: bool = ..., lets_hoisted: bool = ... + ) -> bindings._Expr: ... - def _expr_to_egg(self, expr_decl: ExprDecl, *, expr_to_let: bool = False) -> bindings._Expr: # noqa: PLR0912,C901 + def _expr_to_egg( # noqa: PLR0912,C901 + self, expr_decl: ExprDecl, *, expr_to_let: bool = False, lets_hoisted: bool = False + ) -> bindings._Expr: """ Convert an ExprDecl to an egg expression. + + `lets_hoisted` is passed on to the arguments, see `typed_expr_to_egg`. """ if expr_to_let: try: @@ -1400,7 +1419,7 @@ def _expr_to_egg(self, expr_decl: ExprDecl, *, expr_to_let: bool = False) -> bin res = bindings.Lit(span(), l) case CallDecl() | GetCostDecl(): egg_fn, typed_args = self.translate_call(expr_decl) - egg_args = [self.typed_expr_to_egg(a, expr_to_let) for a in typed_args] + egg_args = [self.typed_expr_to_egg(a, expr_to_let, lets_hoisted=lets_hoisted) for a in typed_args] res = bindings.Call(span(), egg_fn, egg_args) case PyObjectDecl(value): res = bindings.Call( @@ -1415,7 +1434,7 @@ def _expr_to_egg(self, expr_decl: ExprDecl, *, expr_to_let: bool = False) -> bin "unstable-fn", [ bindings.Lit(span(), bindings.String(egg_fn)), - *[self.typed_expr_to_egg(arg, expr_to_let) for arg in typed_args], + *[self.typed_expr_to_egg(arg, expr_to_let, lets_hoisted=lets_hoisted) for arg in typed_args], ], ) case ValueDecl(): diff --git a/python/tests/test_high_level.py b/python/tests/test_high_level.py index f6eb1fcb..dd17d143 100644 --- a/python/tests/test_high_level.py +++ b/python/tests/test_high_level.py @@ -19,6 +19,7 @@ import egglog.bindings as egg_bindings import egglog.builtins as egg_builtins +import egglog.egraph_state as egraph_state_module from egglog import * from egglog.declarations import ( BiRewriteDecl, @@ -1039,6 +1040,49 @@ def test_shared_subexpression_lowering_is_deterministic_across_processes() -> No assert transcripts[1:] == transcripts[:-1] +def test_shared_subexpression_scan_runs_once_per_let(monkeypatch: pytest.MonkeyPatch) -> None: + """ + Converting a deep term scans it for shared subexpressions once, plus once per synthetic let + body, instead of once per subterm (which made lowering quadratic in the depth of the term). + """ + + class ScanLeaf(Expr): + def __init__(self, name: StringLike) -> None: ... + + class ScanList(Expr): + @classmethod + def nil(cls) -> ScanList: ... + + @classmethod + def cons(cls, left: ScanLeaf, right: ScanLeaf, tail: ScanList) -> ScanList: ... + + original_scan = egraph_state_module._exprs_multiple_parents + scans = 0 + + def counting_scan(typed_expr: TypedExprDecl) -> list[TypedExprDecl]: + nonlocal scans + scans += 1 + return original_scan(typed_expr) + + monkeypatch.setattr(egraph_state_module, "_exprs_multiple_parents", counting_scan) + + def scans_to_let(n: int) -> int: + nonlocal scans + shared = ScanLeaf("shared") + term = ScanList.nil() + for i in range(n): + term = ScanList.cons(shared, ScanLeaf(f"x{i}"), term) + egraph = EGraph(save_egglog_string=True) + scans = 0 + egraph.let("root", term) + # Only the shared leaf is hoisted, and every cons cell refers to it. + assert egraph.as_egglog_string.count("(let $__expr_") == 1 + assert egraph.as_egglog_string.count("$__expr_0") == n + 1 + return scans + + assert scans_to_let(100) == scans_to_let(5) + + def test_freeze_omits_synthetic_let_bindings() -> None: class FreezeLetNum(Expr): @classmethod From b0fcdc00f42de72425e0b77b478a8e317bd4fcfe Mon Sep 17 00:00:00 2001 From: lubarsky Date: Thu, 1 Oct 2026 01:26:15 +0300 Subject: [PATCH 2/2] Scan cost lookup arguments for shared subexpressions `_expr_to_egg` converts a cost lookup's arguments with `lets_hoisted`, but `_exprs_multiple_parents` never looked inside a `GetCostDecl`. Since the previous commit, a subexpression shared inside a `get_cost(...)` argument of a scanned expression was therefore written out inline instead of being hoisted into a let (14.0.0 found it by scanning the argument again). The scan now traverses `GetCostDecl.args` like `CallDecl.args`, so it visits every child that is converted with `lets_hoisted`. Found by CodeRabbit's review of #442. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01BnJo4ywpA3FufMmo79Av6f --- python/egglog/egraph_state.py | 11 ++++++++--- python/tests/test_high_level.py | 31 +++++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+), 3 deletions(-) diff --git a/python/egglog/egraph_state.py b/python/egglog/egraph_state.py index c795b924..366b3a0f 100644 --- a/python/egglog/egraph_state.py +++ b/python/egglog/egraph_state.py @@ -1730,7 +1730,12 @@ def _sanitize_egg_ident(input_string: str) -> str: def _exprs_multiple_parents(typed_expr: TypedExprDecl) -> list[TypedExprDecl]: - """Return multiply-parented expressions in deterministic preorder for stable synthetic let names.""" + """ + Return multiply-parented expressions in deterministic preorder for stable synthetic let names. + + Visits every child that `_expr_to_egg` converts with `lets_hoisted`, so the arguments of calls, + cost lookups and partial calls are not scanned again when they are converted. + """ parent_counts: dict[TypedExprDecl, int] = {} traversal_order: list[TypedExprDecl] = [] traversed: set[TypedExprDecl] = set() @@ -1743,7 +1748,7 @@ def _exprs_multiple_parents(typed_expr: TypedExprDecl) -> list[TypedExprDecl]: if node is not typed_expr: traversal_order.append(node) match node.expr: - case CallDecl(args=args) | PartialCallDecl(CallDecl(args=args)): + case CallDecl(args=args) | GetCostDecl(args=args) | PartialCallDecl(CallDecl(args=args)): for child in args: parent_counts[child] = parent_counts.get(child, 0) + 1 stack.extend(reversed(args)) @@ -1765,7 +1770,7 @@ def _contains_unbound_var(typed_expr: TypedExprDecl) -> bool: match node.expr: case UnboundVarDecl(): return True - case CallDecl(args=args) | PartialCallDecl(CallDecl(args=args)): + case CallDecl(args=args) | GetCostDecl(args=args) | PartialCallDecl(CallDecl(args=args)): stack.extend(args) case _: pass diff --git a/python/tests/test_high_level.py b/python/tests/test_high_level.py index dd17d143..b81b2191 100644 --- a/python/tests/test_high_level.py +++ b/python/tests/test_high_level.py @@ -1083,6 +1083,37 @@ def scans_to_let(n: int) -> int: assert scans_to_let(100) == scans_to_let(5) +def test_shared_subexpression_inside_cost_lookup_is_hoisted() -> None: + """ + The scan for shared subexpressions also looks inside a cost lookup's arguments, since they are + converted without scanning again (see `test_shared_subexpression_scan_runs_once_per_let`). + """ + + class CostLeaf(Expr): + def __init__(self, name: StringLike) -> None: ... + + class CostPair(Expr): + def __init__(self, left: CostLeaf, right: CostLeaf) -> None: ... + + class CostRoot(Expr): + def __init__(self, pair: CostPair) -> None: ... + + class CostBox(Expr): + def __init__(self, cost: i64Like) -> None: ... + + shared = CostLeaf("shared") + root = CostRoot(CostPair(shared, shared)) + egraph = EGraph(save_egglog_string=True) + egraph.register(set_cost(root, i64(3))) + start = len(egraph.as_egglog_string) + egraph.register(CostBox(get_cost(root))) + commands = egraph.as_egglog_string[start:] + # The shared leaf is hoisted into one synthetic let, which the pair refers to twice. + assert commands.count("(let $__expr_") == 1 + let_name = commands.split("(let ", 1)[1].split(" ", 1)[0] + assert commands.count(let_name) == 3 + + def test_freeze_omits_synthetic_let_bindings() -> None: class FreezeLetNum(Expr): @classmethod