Skip to content

Commit 1104233

Browse files
Merge pull request #3022 from devitocodes/hotfix-call-from-pointer
compiler: Hotfix CallFromPointer
2 parents 3d35a27 + b7f97d6 commit 1104233

5 files changed

Lines changed: 78 additions & 22 deletions

File tree

‎devito/finite_differences/differentiable.py‎

Lines changed: 27 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -21,8 +21,8 @@
2121
from devito.finite_differences.tools import coeff_priority, make_shift_x0
2222
from devito.logger import warning
2323
from devito.tools import (
24-
as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype, is_integer,
25-
is_number, memoized_func, split
24+
Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype,
25+
is_integer, is_number, memoized_func, split
2626
)
2727
from devito.types import Array, DimensionTuple, Evaluable, StencilDimension
2828
from devito.types.basic import AbstractFunction, Indexed
@@ -34,6 +34,7 @@
3434
'EvalDerivative',
3535
'Imag',
3636
'IndexDerivative',
37+
'IndexDerivativeProperty',
3738
'Real',
3839
'Weights',
3940
]
@@ -1045,14 +1046,23 @@ def value(self, idx):
10451046
return self[idx]
10461047

10471048

1049+
class IndexDerivativeProperty(Tag):
1050+
1051+
"""A property controlling how an `IndexDerivative` is lowered."""
1052+
1053+
10481054
class IndexDerivative(IndexSum):
10491055

10501056
__rargs__ = ('expr', 'mapper')
1051-
__rkwargs__ = IndexSum.__rkwargs__ + ('deriv_order',)
1057+
__rkwargs__ = IndexSum.__rkwargs__ + ('deriv_order', 'properties')
10521058

1053-
def __new__(cls, expr, mapper, deriv_order=None, **kwargs):
1059+
def __new__(cls, expr, mapper, deriv_order=None, properties=(), **kwargs):
10541060
dimensions = as_tuple(set(mapper.values()))
10551061

1062+
properties = frozenset(as_tuple(properties))
1063+
if not all(isinstance(i, IndexDerivativeProperty) for i in properties):
1064+
raise ValueError("Expected IndexDerivative properties")
1065+
10561066
# Detect the Weights among the arguments
10571067
weightss = []
10581068
for a in expr.args:
@@ -1073,20 +1083,25 @@ def __new__(cls, expr, mapper, deriv_order=None, **kwargs):
10731083
obj._mapper = frozendict(mapper)
10741084

10751085
obj._deriv_order = deriv_order
1086+
obj._properties = properties
10761087

10771088
return obj
10781089

10791090
def _hashable_content(self):
1080-
return super()._hashable_content() + (self.mapper,)
1091+
properties = tuple(sorted(map(str, self.properties)))
1092+
return super()._hashable_content() + (self.mapper, properties)
10811093

10821094
def compare(self, other):
10831095
if self is other:
10841096
return 0
10851097
n1 = self.__class__
10861098
n2 = other.__class__
10871099
if n1.__name__ == n2.__name__:
1100+
p1 = tuple(sorted(map(str, self.properties)))
1101+
p2 = tuple(sorted(map(str, other.properties)))
10881102
return (self.weights.compare(other.weights) or
1089-
self.base.compare(other.base))
1103+
self.base.compare(other.base) or
1104+
(p1 > p2) - (p1 < p2))
10901105
else:
10911106
return super().compare(other)
10921107

@@ -1110,6 +1125,10 @@ def mapper(self):
11101125
def deriv_order(self):
11111126
return self._deriv_order
11121127

1128+
@property
1129+
def properties(self):
1130+
return self._properties
1131+
11131132
@property
11141133
def depth(self):
11151134
iderivs = self.expr.find(IndexDerivative)
@@ -1289,7 +1308,8 @@ def _diff2sympy(obj):
12891308
# Handle special objects
12901309
if isinstance(obj, DiffDerivative):
12911310
return IndexDerivative(*args, obj.mapper,
1292-
deriv_order=obj.deriv_order), True
1311+
deriv_order=obj.deriv_order,
1312+
properties=obj.properties), True
12931313

12941314
# Handle generic objects such as arithmetic operations
12951315
try:

‎devito/passes/clusters/derivatives.py‎

Lines changed: 24 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -127,13 +127,15 @@ def _(expr, c, ispace, weights, reusables, mapper, **kwargs):
127127
@_core.register(IndexDerivative)
128128
def _(expr, c, ispace, weights, reusables, mapper, **kwargs):
129129
sregistry = kwargs['sregistry']
130-
options = kwargs['options']
131130

132-
try:
133-
cbk0 = deriv_schedule_registry[options['deriv-schedule']]
134-
cbk1 = deriv_unroll_registry[options['deriv-unroll']]
135-
except KeyError:
136-
raise ValueError("Unknown derivative lowering mode") from None
131+
known = set(deriv_schedule_registry) | set(deriv_unroll_registry)
132+
if not known.issuperset(expr.properties):
133+
raise ValueError("Unknown derivative lowering property")
134+
135+
cbk0 = _select_callback(expr.properties, deriv_schedule_registry,
136+
_lower_index_derivative_base)
137+
cbk1 = _select_callback(expr.properties, deriv_unroll_registry,
138+
_lower_index_derivative_base_unroll)
137139

138140
# Lower the IndexDerivative
139141
init, ideriv = cbk0(expr)
@@ -203,14 +205,24 @@ def _lower_index_derivative_base(ideriv):
203205
return S.Zero, ideriv
204206

205207

206-
deriv_schedule_registry = {
207-
'basic': _lower_index_derivative_base,
208-
}
208+
def _lower_index_derivative_base_unroll(init, ideriv, ispace):
209+
return init, ideriv.expr, ispace
210+
211+
212+
def _select_callback(properties, registry, default):
213+
found = properties.intersection(registry)
214+
if len(found) > 1:
215+
raise ValueError("Incompatible derivative lowering properties")
216+
elif found:
217+
return registry[next(iter(found))]
218+
else:
219+
return default
220+
221+
222+
deriv_schedule_registry = {}
209223

210224

211-
deriv_unroll_registry = {
212-
False: lambda init, ideriv, ispace: (init, ideriv.expr, ispace)
213-
}
225+
deriv_unroll_registry = {}
214226

215227

216228
class CDE(Queue):

‎devito/symbolics/extended_sympy.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -263,6 +263,12 @@ def base(self):
263263
def bound_symbols(self):
264264
return {self.call}
265265

266+
@property
267+
def canonical_variables(self):
268+
# `call` is bound to keep it out of `free_symbols`, but it names a C
269+
# call or member and therefore must not be canonicalized by SymPy
270+
return {}
271+
266272
@property
267273
def free_symbols(self):
268274
return super().free_symbols - self.bound_symbols

‎tests/test_derivatives.py‎

Lines changed: 20 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,8 @@
1010
)
1111
from devito.finite_differences import Derivative, Differentiable, diffify
1212
from devito.finite_differences.differentiable import (
13-
Add, DiffDerivative, EvalDerivative, IndexDerivative, IndexSum, Weights, interp_for_fd
13+
Add, DiffDerivative, EvalDerivative, IndexDerivative, IndexDerivativeProperty,
14+
IndexSum, Weights, interp_for_fd
1415
)
1516
from devito.symbolics import indexify, retrieve_indexed
1617
from devito.types.dimension import StencilDimension
@@ -1084,14 +1085,30 @@ def test_index_derivative(self):
10841085
idxder = IndexDerivative(ui*w, {x: i})
10851086

10861087
assert simplify(idxder.evaluate - (-0.5*u + 0.5*ui.subs(i, 2))) == 0
1088+
assert idxder.properties == frozenset()
1089+
1090+
# Lowering properties are part of the IndexDerivative identity and
1091+
# survive reconstruction
1092+
fold = IndexDerivativeProperty('fold')
1093+
unroll = IndexDerivativeProperty('unroll')
1094+
idxder1 = idxder._rebuild(properties=(fold, unroll))
1095+
assert idxder1.properties == frozenset([fold, unroll])
1096+
assert idxder1 != idxder
1097+
assert len({idxder, idxder1}) == 2
1098+
assert idxder1._rebuild() == idxder1
1099+
assert idxder1.compare(idxder) != 0
1100+
1101+
with pytest.raises(ValueError, match="Expected IndexDerivative properties"):
1102+
idxder._rebuild(properties=('fold', 'unroll'))
10871103

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

10911107
vi0 = v.subs(x, x + i*x.spacing)
1092-
vi1 = idxder.subs(ui, vi0)
1108+
vi1 = idxder1.subs(ui, vi0)
10931109

1094-
assert IndexDerivative(vi0*w, {x: i}) == vi1
1110+
assert IndexDerivative(vi0*w, {x: i},
1111+
properties=idxder1.properties) == vi1
10951112

10961113
def test_dx2(self):
10971114
grid = Grid(shape=(4, 4))

‎tests/test_symbolics.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -306,6 +306,7 @@ def test_field_from_composite():
306306
# Test reconstruction
307307
ffc3 = ffc0.func(*ffc0.args)
308308
assert ffc0 == ffc3
309+
assert ffc1.as_dummy() == ffc1
309310

310311
# Free symbols
311312
assert ffc1.free_symbols == {s}

0 commit comments

Comments
 (0)