diff --git a/devito/ir/iet/visitors.py b/devito/ir/iet/visitors.py index a3d8827c4b..6186ff4e76 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 7ffc461591..59def527f4 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 2df7dc516d..3f17080312 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 37e4ffca81..e5b4617e41 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 b0acb0cdcd..bdbd5dd6da 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)