From eae00397164edee14c794a5722fb71ddb51c7628 Mon Sep 17 00:00:00 2001 From: Fabio Luporini Date: Thu, 1 Oct 2026 11:31:50 +0000 Subject: [PATCH] compiler: Remove IET visitor result memoization Cached results and cache partitions can keep temporary IET roots alive despite weak cache keys. Remove shared result memoization from FindNodes, FindSymbols and FindApplications, along with the unused helper and the redundant FindWithin.visit override. Add regression coverage for root and child retention, scope query targets, returned back-references, repeated visits, subclasses and copied Operators. For a space-order-16 3D TTI Born Operator, six paired runs changed median construction and code generation from 21.18s to 21.24s with identical C. --- devito/ir/iet/visitors.py | 18 +------- devito/tools/memoization.py | 50 ---------------------- tests/test_caching.py | 82 ++++++++++++++++++++++++++++++++++++- tests/test_tools.py | 38 +---------------- tests/test_visitors.py | 39 ++++++++++++++++-- 5 files changed, 119 insertions(+), 108 deletions(-) diff --git a/devito/ir/iet/visitors.py b/devito/ir/iet/visitors.py index a3d8827c4ba..6186ff4e763 100644 --- a/devito/ir/iet/visitors.py +++ b/devito/ir/iet/visitors.py @@ -28,7 +28,7 @@ from devito.symbolics.extended_dtypes import NoDeclStruct from devito.tools import ( GenericVisitor, as_tuple, c_restrict_void_p, filter_ordered, filter_sorted, flatten, - is_external_ctype, memoized_weak_meth, natural_sort_key + is_external_ctype, natural_sort_key ) from devito.types import ( ArrayObject, CompositeObject, DeviceMap, Dimension, IndexedData, Pointer @@ -1127,10 +1127,6 @@ def __init__(self, mode: str = 'symbolics') -> None: else: self.rule = lambda n: chain(*[self.rules[mode](n) for mode in modes]) - @memoized_weak_meth(key=lambda i: i.mode, freeze=tuple, thaw=list) - def visit(self, o, *args, **kwargs): - return super().visit(o, *args, **kwargs) - def _post_visit(self, ret): return sorted(filter_ordered(ret, key=id), key=natural_sort_key) @@ -1180,10 +1176,6 @@ def __init__(self, match: type, mode: str = 'type') -> None: self.mode = mode self.rule = self.rules[mode] - @memoized_weak_meth(key=lambda i: (i.match, i.mode), freeze=tuple, thaw=list) - def visit(self, o, *args, **kwargs): - return super().visit(o, *args, **kwargs) - def visit_Node(self, o: Node, **kwargs) -> Iterator[Node]: if self.rule(self.match, o): yield o @@ -1204,10 +1196,6 @@ def __init__(self, match: type, start: Node, stop: Node | None = None) -> None: self.start = start self.stop = stop - def visit(self, o, *args, **kwargs): - # `start` and `stop` are part of this visitor's state. - return GenericVisitor.visit(self, o, *args, **kwargs) - def visit_object(self, o: object, flag: bool = False) -> LazyVisit[Node, bool]: yield from () return flag # noqa: B901 @@ -1258,10 +1246,6 @@ def __init__(self, cls: type[ApplicationType] = Application): self.cls = cls self.match = lambda i: isinstance(i, cls) and not isinstance(i, Basic) - @memoized_weak_meth(key=lambda i: i.cls, freeze=frozenset, thaw=set) - def visit(self, o, *args, **kwargs): - return super().visit(o, *args, **kwargs) - def _post_visit(self, ret): return set(ret) diff --git a/devito/tools/memoization.py b/devito/tools/memoization.py index 7ffc461591a..59def527f45 100644 --- a/devito/tools/memoization.py +++ b/devito/tools/memoization.py @@ -2,7 +2,6 @@ from functools import lru_cache, partial, wraps from itertools import tee from typing import TypeVar -from weakref import WeakKeyDictionary __all__ = [ 'CacheInstances', @@ -10,7 +9,6 @@ 'memoized_func', 'memoized_generator', 'memoized_meth', - 'memoized_weak_meth', 'reuse_if_unchanged' ] @@ -191,54 +189,6 @@ def __call__(self, *args, **kwargs): return result -def memoized_weak_meth(*, key=None, freeze=None, thaw=None): - """ - Cache a method result against its first argument using weak references. - - This is useful for visitors operating on temporary IR roots: the cache can - be shared across short-lived visitor instances without keeping those roots - alive. Only calls without extra arguments are cached; all other calls fall - back to the wrapped method. - - Parameters - ---------- - key : callable, optional - A callable receiving ``self`` and returning a hashable cache partition. - freeze : callable, optional - Convert the method result before storing it in the cache. - thaw : callable, optional - Convert the cached value before returning it to the caller. - """ - def decorator(func): - caches = {} - - @wraps(func) - def wrapper(self, o, *args, **kwargs): - if args or kwargs: - return func(self, o, *args, **kwargs) - - try: - partition = key(self) if key is not None else None - cache = caches.setdefault(partition, WeakKeyDictionary()) - ret = cache[o] - except KeyError: - ret = func(self, o) - if freeze is not None: - ret = freeze(ret) - cache[o] = ret - except TypeError: - return func(self, o) - - if thaw is not None: - return thaw(ret) - - return ret - - return wrapper - - return decorator - - # Describes the type of a subclass of CacheInstances InstanceType = TypeVar('InstanceType', bound='CacheInstances', covariant=True) diff --git a/tests/test_caching.py b/tests/test_caching.py index 2df7dc516d8..3f170803123 100644 --- a/tests/test_caching.py +++ b/tests/test_caching.py @@ -1,4 +1,7 @@ +import gc +import pickle import weakref +from copy import copy, deepcopy from ctypes import byref, c_void_p import numpy as np @@ -6,11 +9,12 @@ from sympy import Expr from devito import ( - ConditionalDimension, Constant, DefaultDimension, Dimension, Eq, Function, Grid, + ConditionalDimension, Constant, DefaultDimension, Dimension, Eq, Function, Grid, Min, Operator, SparseFunction, SparseTimeFunction, SubDimension, TensorFunction, TensorTimeFunction, TimeFunction, VectorFunction, VectorTimeFunction, _SymbolCache, clear_cache, solve, switchconfig ) +from devito.ir.iet import Call, FindApplications, FindNodes, FindSymbols, List, Node from devito.types import ( DeviceID, LocalObject, NPThreads, NThreadsBase, Object, Scalar, Symbol, ThreadID ) @@ -790,6 +794,79 @@ class TestMemoryLeaks: Tests ensuring there are no memory leaks. """ + @pytest.mark.parametrize('nested', [False, True]) + def test_findnodes_leakage(self, nested): + """A traversal containing its root must not keep the tree alive.""" + def visit_temporary_tree(): + tree = List(body=[Call('foo')] if nested else []) + nodes = [tree, *tree.body] + references = [weakref.ref(i) for i in nodes] + assert FindNodes(Node).visit(tree) == nodes + return references + + references = visit_temporary_tree() + gc.collect() + clear_cache() + + assert all(i() is None for i in references) + + @pytest.mark.parametrize('kind', ['list', 'operator', 'jitted-operator']) + @pytest.mark.parametrize('copier', [ + copy, deepcopy, pytest.param(lambda o: pickle.loads(pickle.dumps(o)), id='pickle') + ]) + def test_copied_iet_leakage(self, kind, copier): + """An unvisited copy must not retain the original traversal's root.""" + def copy_temporary_tree(): + if kind != 'list': + f = Function(name='f', grid=Grid(shape=(3, 3))) + tree = Operator(Eq(f, f + 1)) + if kind == 'jitted-operator': + tree.apply() + else: + tree = List() + assert FindNodes(Node).visit(tree)[0] is tree + return copier(tree), weakref.ref(tree) + + copied, reference = copy_temporary_tree() + clear_cache() + assert reference() is None + assert FindNodes(Node).visit(copied)[0] is copied + + @pytest.mark.parametrize('match_root', [False, True]) + def test_findnodes_scope_leakage(self, match_root): + """A scope query must not keep its target or containing tree alive.""" + def visit_temporary_tree(): + child = Call('foo') + tree = List(body=[child]) + match = tree if match_root else child + expected = [] if match_root else [tree] + assert FindNodes(match, mode='scope').visit(tree) == expected + return weakref.ref(tree), weakref.ref(child) + + references = visit_temporary_tree() + clear_cache() + assert all(i() is None for i in references) + + @pytest.mark.parametrize('visitor', [FindSymbols, FindApplications]) + def test_visitor_leakage_backref(self, visitor): + """A returned object's reference to the root must not cause a leak.""" + def visit_temporary_tree(): + if visitor is FindSymbols: + result = Object(name='context', dtype=c_void_p) + tree = Call('foo', arguments=[result]) + result.value = tree + else: + result = Min(Symbol(name='s'), 1) + tree = Call('foo', arguments=[result]) + result._test_owner = tree + + assert result in visitor().visit(tree) + return weakref.ref(tree), weakref.ref(result) + + references = visit_temporary_tree() + clear_cache() + assert all(i() is None for i in references) + def test_operator_leakage_function(self): """ Test to ensure that Operator creation does not cause memory leaks for @@ -806,6 +883,9 @@ def test_operator_leakage_function(self): # Create operator and delete everything again op = Operator(Eq(f, 2 * g)) w_op = weakref.ref(op) + FindNodes(Node).visit(op) + FindSymbols().visit(op) + FindApplications().visit(op) del op del f del g diff --git a/tests/test_tools.py b/tests/test_tools.py index 37e4ffca81e..e5b4617e411 100644 --- a/tests/test_tools.py +++ b/tests/test_tools.py @@ -8,7 +8,7 @@ from devito import Eq, Operator, switchenv from devito.tools import ( CacheInstances, DefaultFrozenDict, UnboundedMultiTuple, UnboundTuple, ctypes_to_cstr, - filter_ordered, memoized_meth, memoized_weak_meth, toposort, transitive_closure + filter_ordered, memoized_meth, toposort, transitive_closure ) from devito.types.basic import Symbol @@ -88,42 +88,6 @@ def f(self, x=None): assert obj.calls == 4 -def test_memoized_weak_meth(): - - class Root: - pass - - class Obj: - - def __init__(self, mode): - self.mode = mode - self.calls = 0 - - @memoized_weak_meth(key=lambda i: i.mode, freeze=tuple, thaw=list) - def f(self, root): - self.calls += 1 - return [self.mode] - - root = Root() - obj0 = Obj('a') - obj1 = Obj('a') - obj2 = Obj('b') - - ret = obj0.f(root) - ret.append('mutated') - - assert obj1.f(root) == ['a'] - assert obj0.calls == 1 - assert obj1.calls == 0 - - assert obj2.f(root) == ['b'] - assert obj2.calls == 1 - - assert obj0.f([]) == ['a'] - assert obj0.f([]) == ['a'] - assert obj0.calls == 3 - - def test_default_frozen_dict(): mapper = DefaultFrozenDict({'a': 'b'}, default='c') diff --git a/tests/test_visitors.py b/tests/test_visitors.py index b0acb0cdcd2..bdbd5dd6dad 100644 --- a/tests/test_visitors.py +++ b/tests/test_visitors.py @@ -8,8 +8,8 @@ from devito.ir.equations import DummyEq from devito.ir.iet import ( Block, Call, Callable, Conditional, Definition, Expression, FindApplications, - FindNodes, FindSections, FindSymbols, FindWithin, IsPerfectIteration, Iteration, - MapNodes, Transformer, Uxreplace, printAST + FindNodes, FindSections, FindSymbols, FindWithin, IsPerfectIteration, Iteration, List, + MapNodes, Node, Transformer, Uxreplace, printAST ) from devito.symbolics import ListInitializer from devito.types import Array, LocalObject, SpaceDimension, Symbol @@ -213,7 +213,40 @@ def test_find_sections(exprs, block1, block2, block3): assert len(found[2]) == 1 -def test_find_within_not_cached_like_findnodes(block3): +@pytest.mark.parametrize('match', [Node, List, Call]) +def test_find_nodes_repeated(match): + call0 = Call('foo') + call1 = Call('bar') + inner = List(body=[call1]) + tree = List(body=[call0, inner]) + expected = [i for i in [tree, call0, inner, call1] if isinstance(i, match)] + + finder = FindNodes(match) + result = finder.visit(tree) + assert result == expected + result.clear() + assert finder.visit(tree) == expected + assert FindNodes(match).visit(tree) == expected + + +def test_find_nodes_subclass(): + + class CallsOnly(FindNodes): + + def visit_Node(self, o, **kwargs): + if isinstance(o, Call): + yield o + for child in o.children: + yield from self._visit(child, **kwargs) + + call = Call('foo') + tree = List(body=[call]) + assert FindNodes(Node).visit(tree) == [tree, call] + assert CallsOnly(Node).visit(tree) == [call] + assert FindNodes(Node).visit(tree) == [tree, call] + + +def test_find_within_bounds(block3): expr0 = FindWithin(Expression, block3.nodes[0], block3.nodes[1]).visit(block3) expr1 = FindWithin(Expression, block3.nodes[1], block3.nodes[2]).visit(block3)