Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 1 addition & 17 deletions devito/ir/iet/visitors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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)

Expand Down
50 changes: 0 additions & 50 deletions devito/tools/memoization.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,13 @@
from functools import lru_cache, partial, wraps
from itertools import tee
from typing import TypeVar
from weakref import WeakKeyDictionary

__all__ = [
'CacheInstances',
'cached_hash',
'memoized_func',
'memoized_generator',
'memoized_meth',
'memoized_weak_meth',
'reuse_if_unchanged'
]

Expand Down Expand Up @@ -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)

Expand Down
82 changes: 81 additions & 1 deletion tests/test_caching.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,20 @@
import gc
import pickle
import weakref
from copy import copy, deepcopy
from ctypes import byref, c_void_p

import numpy as np
import pytest
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
)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
38 changes: 1 addition & 37 deletions tests/test_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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')

Expand Down
39 changes: 36 additions & 3 deletions tests/test_visitors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
Loading