Skip to content
Open
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
34 changes: 27 additions & 7 deletions devito/finite_differences/differentiable.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,8 +21,8 @@
from devito.finite_differences.tools import coeff_priority, make_shift_x0
from devito.logger import warning
from devito.tools import (
as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype, is_integer,
is_number, memoized_func, split
Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype,
is_integer, is_number, memoized_func, split
)
from devito.types import Array, DimensionTuple, Evaluable, StencilDimension
from devito.types.basic import AbstractFunction, Indexed
Expand All @@ -34,6 +34,7 @@
'EvalDerivative',
'Imag',
'IndexDerivative',
'IndexDerivativeProperty',
'Real',
'Weights',
]
Expand Down Expand Up @@ -1045,14 +1046,23 @@ def value(self, idx):
return self[idx]


class IndexDerivativeProperty(Tag):

"""A property controlling how an `IndexDerivative` is lowered."""


class IndexDerivative(IndexSum):

__rargs__ = ('expr', 'mapper')
__rkwargs__ = IndexSum.__rkwargs__ + ('deriv_order',)
__rkwargs__ = IndexSum.__rkwargs__ + ('deriv_order', 'properties')

def __new__(cls, expr, mapper, deriv_order=None, **kwargs):
def __new__(cls, expr, mapper, deriv_order=None, properties=(), **kwargs):
dimensions = as_tuple(set(mapper.values()))

properties = frozenset(as_tuple(properties))
if not all(isinstance(i, IndexDerivativeProperty) for i in properties):
raise ValueError("Expected IndexDerivative properties")

# Detect the Weights among the arguments
weightss = []
for a in expr.args:
Expand All @@ -1073,20 +1083,25 @@ def __new__(cls, expr, mapper, deriv_order=None, **kwargs):
obj._mapper = frozendict(mapper)

obj._deriv_order = deriv_order
obj._properties = properties

return obj

def _hashable_content(self):
return super()._hashable_content() + (self.mapper,)
properties = tuple(sorted(map(str, self.properties)))
return super()._hashable_content() + (self.mapper, properties)

def compare(self, other):
if self is other:
return 0
n1 = self.__class__
n2 = other.__class__
if n1.__name__ == n2.__name__:
p1 = tuple(sorted(map(str, self.properties)))
p2 = tuple(sorted(map(str, other.properties)))
return (self.weights.compare(other.weights) or
self.base.compare(other.base))
self.base.compare(other.base) or
(p1 > p2) - (p1 < p2))
else:
return super().compare(other)

Expand All @@ -1110,6 +1125,10 @@ def mapper(self):
def deriv_order(self):
return self._deriv_order

@property
def properties(self):
return self._properties

@property
def depth(self):
iderivs = self.expr.find(IndexDerivative)
Expand Down Expand Up @@ -1289,7 +1308,8 @@ def _diff2sympy(obj):
# Handle special objects
if isinstance(obj, DiffDerivative):
return IndexDerivative(*args, obj.mapper,
deriv_order=obj.deriv_order), True
deriv_order=obj.deriv_order,
properties=obj.properties), True

# Handle generic objects such as arithmetic operations
try:
Expand Down
36 changes: 24 additions & 12 deletions devito/passes/clusters/derivatives.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,13 +127,15 @@ def _(expr, c, ispace, weights, reusables, mapper, **kwargs):
@_core.register(IndexDerivative)
def _(expr, c, ispace, weights, reusables, mapper, **kwargs):
sregistry = kwargs['sregistry']
options = kwargs['options']

try:
cbk0 = deriv_schedule_registry[options['deriv-schedule']]
cbk1 = deriv_unroll_registry[options['deriv-unroll']]
except KeyError:
raise ValueError("Unknown derivative lowering mode") from None
known = set(deriv_schedule_registry) | set(deriv_unroll_registry)
if not known.issuperset(expr.properties):
raise ValueError("Unknown derivative lowering property")

cbk0 = _select_callback(expr.properties, deriv_schedule_registry,
_lower_index_derivative_base)
cbk1 = _select_callback(expr.properties, deriv_unroll_registry,
_lower_index_derivative_base_unroll)

# Lower the IndexDerivative
init, ideriv = cbk0(expr)
Expand Down Expand Up @@ -203,14 +205,24 @@ def _lower_index_derivative_base(ideriv):
return S.Zero, ideriv


deriv_schedule_registry = {
'basic': _lower_index_derivative_base,
}
def _lower_index_derivative_base_unroll(init, ideriv, ispace):
return init, ideriv.expr, ispace
Comment on lines +208 to +209

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why does this need to take init and ispace as args?



def _select_callback(properties, registry, default):
found = properties.intersection(registry)
if len(found) > 1:
raise ValueError("Incompatible derivative lowering properties")
elif found:
return registry[next(iter(found))]
else:
return default
Comment on lines +214 to +219

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

if not found:
    return default

if len(found) > 1:
    raise ValueError("Incompatible derivative lowering properties")

return registry[next(iter(found))]

would probably match our codestyle more?



deriv_schedule_registry = {}


deriv_unroll_registry = {
False: lambda init, ideriv, ispace: (init, ideriv.expr, ispace)
}
deriv_unroll_registry = {}


class CDE(Queue):
Expand Down
6 changes: 6 additions & 0 deletions devito/symbolics/extended_sympy.py
Original file line number Diff line number Diff line change
Expand Up @@ -263,6 +263,12 @@ def base(self):
def bound_symbols(self):
return {self.call}

@property
def canonical_variables(self):
# `call` is bound to keep it out of `free_symbols`, but it names a C
# call or member and therefore must not be canonicalized by SymPy
return {}

@property
def free_symbols(self):
return super().free_symbols - self.bound_symbols
Expand Down
23 changes: 20 additions & 3 deletions tests/test_derivatives.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@
)
from devito.finite_differences import Derivative, Differentiable, diffify
from devito.finite_differences.differentiable import (
Add, DiffDerivative, EvalDerivative, IndexDerivative, IndexSum, Weights, interp_for_fd
Add, DiffDerivative, EvalDerivative, IndexDerivative, IndexDerivativeProperty,
IndexSum, Weights, interp_for_fd
)
from devito.symbolics import indexify, retrieve_indexed
from devito.types.dimension import StencilDimension
Expand Down Expand Up @@ -1084,14 +1085,30 @@ def test_index_derivative(self):
idxder = IndexDerivative(ui*w, {x: i})

assert simplify(idxder.evaluate - (-0.5*u + 0.5*ui.subs(i, 2))) == 0
assert idxder.properties == frozenset()

# Lowering properties are part of the IndexDerivative identity and
# survive reconstruction
fold = IndexDerivativeProperty('fold')
unroll = IndexDerivativeProperty('unroll')
idxder1 = idxder._rebuild(properties=(fold, unroll))
assert idxder1.properties == frozenset([fold, unroll])
assert idxder1 != idxder
assert len({idxder, idxder1}) == 2
assert idxder1._rebuild() == idxder1
assert idxder1.compare(idxder) != 0

with pytest.raises(ValueError, match="Expected IndexDerivative properties"):
idxder._rebuild(properties=('fold', 'unroll'))

# Make sure subs works as expected
v = Function(name="v", grid=grid, space_order=so)

vi0 = v.subs(x, x + i*x.spacing)
vi1 = idxder.subs(ui, vi0)
vi1 = idxder1.subs(ui, vi0)

assert IndexDerivative(vi0*w, {x: i}) == vi1
assert IndexDerivative(vi0*w, {x: i},
properties=idxder1.properties) == vi1

def test_dx2(self):
grid = Grid(shape=(4, 4))
Expand Down
1 change: 1 addition & 0 deletions tests/test_symbolics.py
Original file line number Diff line number Diff line change
Expand Up @@ -306,6 +306,7 @@ def test_field_from_composite():
# Test reconstruction
ffc3 = ffc0.func(*ffc0.args)
assert ffc0 == ffc3
assert ffc1.as_dummy() == ffc1

# Free symbols
assert ffc1.free_symbols == {s}
Expand Down
Loading