Skip to content

Commit e93685d

Browse files
committed
compiler: Attach staggering metadata to IndexDerivative
1 parent 33c3b12 commit e93685d

4 files changed

Lines changed: 209 additions & 14 deletions

File tree

‎devito/finite_differences/differentiable.py‎

Lines changed: 29 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1048,9 +1048,9 @@ def value(self, idx):
10481048
class IndexDerivative(IndexSum):
10491049

10501050
__rargs__ = ('expr', 'mapper')
1051-
__rkwargs__ = IndexSum.__rkwargs__ + ('deriv_order',)
1051+
__rkwargs__ = IndexSum.__rkwargs__ + ('deriv_order', 'staggering')
10521052

1053-
def __new__(cls, expr, mapper, deriv_order=None, **kwargs):
1053+
def __new__(cls, expr, mapper, deriv_order=None, staggering=None, **kwargs):
10541054
dimensions = as_tuple(set(mapper.values()))
10551055

10561056
# Detect the Weights among the arguments
@@ -1073,11 +1073,20 @@ def __new__(cls, expr, mapper, deriv_order=None, **kwargs):
10731073
obj._mapper = frozendict(mapper)
10741074

10751075
obj._deriv_order = deriv_order
1076+
obj._staggering = staggering
10761077

10771078
return obj
10781079

1080+
@cached_property
1081+
def _metadata(self):
1082+
# SymPy's canonical sorting also compares the hashable content directly.
1083+
# Use comparable objects, including empty tuples for unknown metadata
1084+
return (sympy.Dict(*self.mapper.items()),
1085+
sympy.Tuple(*as_tuple(self.deriv_order)),
1086+
sympy.Tuple(*as_tuple(self.staggering)))
1087+
10791088
def _hashable_content(self):
1080-
return super()._hashable_content() + (self.mapper,)
1089+
return super()._hashable_content() + self._metadata
10811090

10821091
def compare(self, other):
10831092
if self is other:
@@ -1086,7 +1095,8 @@ def compare(self, other):
10861095
n2 = other.__class__
10871096
if n1.__name__ == n2.__name__:
10881097
return (self.weights.compare(other.weights) or
1089-
self.base.compare(other.base))
1098+
self.base.compare(other.base) or
1099+
super().compare(other))
10901100
else:
10911101
return super().compare(other)
10921102

@@ -1110,6 +1120,19 @@ def mapper(self):
11101120
def deriv_order(self):
11111121
return self._deriv_order
11121122

1123+
@property
1124+
def staggering(self):
1125+
"""
1126+
The requested evaluation staggering relative to the input lattice:
1127+
`centered` on that lattice, `staggered` halfway between its points,
1128+
or None for other or unknown evaluation locations.
1129+
1130+
This classification is independent of differential order, stencil bias,
1131+
and transposition. In particular, `centered` does not imply symmetric
1132+
weights, and interpolation is identified separately by `deriv_order == 0`.
1133+
"""
1134+
return self._staggering
1135+
11131136
@property
11141137
def depth(self):
11151138
iderivs = self.expr.find(IndexDerivative)
@@ -1288,8 +1311,8 @@ def _diff2sympy(obj):
12881311

12891312
# Handle special objects
12901313
if isinstance(obj, DiffDerivative):
1291-
return IndexDerivative(*args, obj.mapper,
1292-
deriv_order=obj.deriv_order), True
1314+
kwargs = {i: getattr(obj, i) for i in obj.__rkwargs__}
1315+
return IndexDerivative(*args, obj.mapper, **kwargs), True
12931316

12941317
# Handle generic objects such as arithmetic operations
12951318
try:

‎devito/finite_differences/finite_difference.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -236,7 +236,8 @@ def make_derivative(expr, dim, fd_order, deriv_order, side, matvec, x0, coeffici
236236
expr = expr._evaluate(expand=False)
237237

238238
deriv = DiffDerivative(
239-
expr*weights, {dim: indices.free_dim}, deriv_order=deriv_order
239+
expr*weights, {dim: indices.free_dim}, deriv_order=deriv_order,
240+
staggering=indices.staggering
240241
)
241242
else:
242243
terms = []

‎devito/finite_differences/tools.py‎

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -146,9 +146,13 @@ class IndexSet(tuple):
146146

147147
"""
148148
The points of a finite-difference expansion.
149+
150+
`staggering` records the scheme's requested evaluation staggering relative
151+
to the input lattice: `centered`, `staggered`, or None if unknown. Index
152+
coordinate changes preserve this classification.
149153
"""
150154

151-
def __new__(cls, dim, indices=None, expr=None, fd=None):
155+
def __new__(cls, dim, indices=None, expr=None, fd=None, staggering=None):
152156
assert indices is not None or expr is not None
153157

154158
if fd is None:
@@ -167,6 +171,7 @@ def __new__(cls, dim, indices=None, expr=None, fd=None):
167171
obj.dim = dim
168172
obj.expr = expr
169173
obj.free_dim = fd
174+
obj.staggering = staggering
170175

171176
return obj
172177

@@ -203,7 +208,8 @@ def transpose(self):
203208
except AttributeError:
204209
expr = None
205210

206-
return IndexSet(self.dim, indices, expr=expr, fd=free_dim)
211+
return IndexSet(self.dim, indices, expr=expr, fd=free_dim,
212+
staggering=self.staggering)
207213

208214
def shift(self, v):
209215
"""
@@ -216,7 +222,8 @@ def shift(self, v):
216222
except TypeError:
217223
expr = None
218224

219-
return IndexSet(self.dim, indices, expr=expr, fd=self.free_dim)
225+
return IndexSet(self.dim, indices, expr=expr, fd=self.free_dim,
226+
staggering=self.staggering)
220227

221228

222229
def make_stencil_dimension(expr, _min, _max):
@@ -287,6 +294,12 @@ def generate_indices(expr, dim, order, side=None, matvec=None, x0=None, nweights
287294

288295
# Evaluation point relative to the expression's grid
289296
mid = (x0 - expr.indices_ref[dim]).subs({dim: 0, dim.spacing: 1})
297+
if (mid % 1).is_zero:
298+
staggering = 'centered'
299+
elif ((mid - S.Half) % 1).is_zero:
300+
staggering = 'staggered'
301+
else:
302+
staggering = None
290303

291304
# Shift for side
292305
side = side or centered
@@ -305,7 +318,7 @@ def generate_indices(expr, dim, order, side=None, matvec=None, x0=None, nweights
305318
d = make_stencil_dimension(expr, o_min, o_max)
306319
iexpr = expr.indices_ref[dim] + d * dim.spacing
307320

308-
return IndexSet(dim, expr=iexpr), x0
321+
return IndexSet(dim, expr=iexpr, staggering=staggering), x0
309322

310323

311324
def make_shift_x0(shift, ndim):

‎tests/test_derivatives.py‎

Lines changed: 161 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
import numpy as np
22
import pytest
3-
from sympy import Float, Symbol, diff, simplify, sympify
3+
from sympy import Float, S, Symbol, diff, simplify, sympify
44

55
from conftest import assert_structure
66
from devito import (
@@ -10,9 +10,12 @@
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, IndexSum, Weights, diff2sympy,
14+
interp_for_fd
1415
)
15-
from devito.symbolics import indexify, retrieve_indexed
16+
from devito.finite_differences.tools import generate_indices
17+
from devito.ir.equations.algorithms import lower_exprs
18+
from devito.symbolics import indexify, retrieve_indexed, search, uxreplace
1619
from devito.types.dimension import StencilDimension
1720
from devito.warnings import DevitoWarning
1821

@@ -1093,6 +1096,161 @@ def test_index_derivative(self):
10931096

10941097
assert IndexDerivative(vi0*w, {x: i}) == vi1
10951098

1099+
@pytest.mark.parametrize('staggered', [None, x, y])
1100+
@pytest.mark.parametrize('deriv_order,fd_order', [
1101+
(0, 4), (0, 16), (1, 2), (1, 4), (1, 16), (2, 4), (2, 16)
1102+
])
1103+
@pytest.mark.parametrize('offset,staggering', [
1104+
(0, 'centered'), (S.Half, 'staggered'), (-S.Half, 'staggered'),
1105+
(1, 'centered'), (S.One/4, None)
1106+
])
1107+
def test_index_derivative_staggering(self, staggered, deriv_order, fd_order,
1108+
offset, staggering):
1109+
grid = Grid(shape=(10, 10))
1110+
x, y = grid.dimensions
1111+
staggered = staggered(grid) if staggered else NODE
1112+
f = Function(name='f', grid=grid, space_order=16, staggered=staggered)
1113+
x0 = {x: f.indices_ref[x] + offset*x.spacing}
1114+
deriv = f.diff(x, deriv_order=deriv_order, fd_order=fd_order, x0=x0)
1115+
1116+
evaluated = deriv._evaluate(expand=False)
1117+
if deriv_order == 0 and offset == 0:
1118+
assert evaluated == f
1119+
return
1120+
1121+
lowered = lower_exprs(diff2sympy(evaluated))
1122+
for i in (evaluated, lowered):
1123+
assert isinstance(i, IndexDerivative)
1124+
assert i.staggering == staggering
1125+
assert i.deriv_order == deriv_order
1126+
1127+
# Classifying the staggering must not change the discrete stencil
1128+
sd, = evaluated.dimensions
1129+
terms = [w*evaluated.base.subs(sd, i) for w, i in
1130+
zip(evaluated.weights.function.weights, sd.range, strict=True)]
1131+
assert simplify(sum(terms) - deriv.evaluate) == 0
1132+
1133+
@pytest.mark.parametrize('offsets,staggerings,fd_order,weights', [
1134+
((0, S.Half), ('centered', 'staggered'), 2, None),
1135+
((S.Half, S.One/4), ('staggered', None), 4,
1136+
[S.One/24, -S(9)/8, S(9)/8, -S.One/24])
1137+
])
1138+
@pytest.mark.parametrize('lowered', [False, True])
1139+
def test_index_derivative_staggering_identity(self, offsets, staggerings, fd_order,
1140+
weights, lowered):
1141+
grid = Grid(shape=(10,))
1142+
x, = grid.dimensions
1143+
f = Function(name='f', grid=grid, space_order=4)
1144+
derivs = [f.dx(fd_order=fd_order, x0={x: x + i*x.spacing}, weights=weights)
1145+
for i in offsets]
1146+
derivs = [i._evaluate(expand=False) for i in derivs]
1147+
if lowered:
1148+
derivs = [lower_exprs(diff2sympy(i)) for i in derivs]
1149+
a, b = derivs
1150+
1151+
# Distinct staggerings can generate exactly the same discrete stencil
1152+
assert tuple(i.staggering for i in derivs) == staggerings
1153+
assert a.args == b.args
1154+
assert a.mapper == b.mapper
1155+
assert a.weights.function.weights == b.weights.function.weights
1156+
assert a != b
1157+
assert len({a, b}) == 2
1158+
assert a.compare(b) == -b.compare(a) != 0
1159+
assert set((a + b).args) == {a, b}
1160+
assert a.evaluate == b.evaluate
1161+
1162+
@pytest.mark.parametrize('offset,staggering', [
1163+
(-S.Half, 'staggered'), (0, 'centered')
1164+
])
1165+
@pytest.mark.parametrize('side', [None, centered, left, right])
1166+
@pytest.mark.parametrize('transpose', [False, True])
1167+
def test_index_derivative_staggering_transpose(self, offset, staggering, side,
1168+
transpose):
1169+
grid = Grid(shape=(10,))
1170+
x, = grid.dimensions
1171+
f = Function(name='f', grid=grid, space_order=4, staggered=x)
1172+
deriv = f.dx(x0={x: f.indices_ref[x] + offset*x.spacing}, side=side)
1173+
if transpose:
1174+
deriv = deriv.T
1175+
evaluated = deriv._evaluate(expand=False)
1176+
assert evaluated.staggering == staggering
1177+
assert lower_exprs(diff2sympy(evaluated)).staggering == staggering
1178+
assert simplify(evaluated.evaluate - deriv.evaluate) == 0
1179+
1180+
@pytest.mark.parametrize('dim', [x, y])
1181+
@pytest.mark.parametrize('offset,staggering', [
1182+
(-S.Half, 'staggered'), (0, 'centered')
1183+
])
1184+
def test_index_derivative_staggering_nested(self, dim, offset, staggering):
1185+
grid = Grid(shape=(10, 10))
1186+
x, y = grid.dimensions
1187+
d = dim(grid)
1188+
f = Function(name='f', grid=grid, space_order=4)
1189+
deriv = f.dx(x0={x: x + x.spacing/2})
1190+
deriv = deriv.diff(d, deriv_order=2, x0={d: d + offset*d.spacing})
1191+
1192+
evaluated = deriv._evaluate(expand=False)
1193+
lowered = lower_exprs(diff2sympy(evaluated))
1194+
for expr in (evaluated, lowered):
1195+
iderivs = search(expr, IndexDerivative)
1196+
assert len(iderivs) == 2
1197+
assert {(i.deriv_order, i.staggering) for i in iderivs} == {
1198+
(1, 'staggered'), (2, staggering)
1199+
}
1200+
1201+
@pytest.mark.parametrize('lowered', [False, True])
1202+
def test_index_derivative_staggering_rebuild(self, lowered):
1203+
grid = Grid(shape=(10,))
1204+
x, = grid.dimensions
1205+
f = Function(name='f', grid=grid, space_order=4)
1206+
g = Function(name='g', grid=grid, space_order=4)
1207+
expr = f.dx(x0={x: x + x.spacing/2})._evaluate(expand=False)
1208+
sd, = expr.dimensions
1209+
base = g.subs(x, x + sd*x.spacing)
1210+
if lowered:
1211+
expr = lower_exprs(diff2sympy(expr))
1212+
base = lower_exprs(base)
1213+
1214+
for rebuilt in (expr.func(base*expr.weights), expr.subs(expr.base, base),
1215+
expr.xreplace({expr.base: base}),
1216+
uxreplace(expr, {expr.base: base})):
1217+
assert rebuilt.base == base
1218+
assert rebuilt.staggering == 'staggered'
1219+
assert rebuilt.deriv_order == 1
1220+
1221+
# Missing metadata must stay distinct from a known centered interpolation
1222+
unknown = expr._rebuild(staggering=None, deriv_order=None)
1223+
on_grid = expr._rebuild(staggering='centered', deriv_order=0)
1224+
assert unknown.staggering is None
1225+
assert len({expr, unknown, on_grid}) == 3
1226+
assert unknown.compare(on_grid) == -on_grid.compare(unknown) != 0
1227+
assert set((expr + unknown + on_grid).args) == {expr, unknown, on_grid}
1228+
1229+
def test_index_derivative_staggering_unknown(self):
1230+
grid = Grid(shape=(10, 10))
1231+
x, y = grid.dimensions
1232+
f = Function(name='f', grid=grid, space_order=4)
1233+
exprs = (f.dx45, f.dx(x0={x: 1}))
1234+
for expr in exprs:
1235+
lowered = lower_exprs(diff2sympy(expr._evaluate(expand=False)))
1236+
iderivs = search(lowered, IndexDerivative)
1237+
assert iderivs
1238+
assert all(i.staggering is None for i in iderivs)
1239+
1240+
@pytest.mark.parametrize('offset,staggering', [
1241+
(0, 'centered'), (1.0, 'centered'), (-1.0, 'centered'),
1242+
(0.5, 'staggered'), (-1.5, 'staggered'), (0.25, None), (0.50000001, None)
1243+
])
1244+
def test_index_set_staggering(self, offset, staggering):
1245+
grid = Grid(shape=(10,))
1246+
x, = grid.dimensions
1247+
f = Function(name='f', grid=grid, space_order=4, staggered=x)
1248+
indices, _ = generate_indices(f, x, 4,
1249+
x0={x: f.indices_ref[x] + offset*x.spacing})
1250+
for i in (indices, indices.transpose(), indices.shift(-x.spacing/2),
1251+
indices.transpose().shift(-x.spacing/2)):
1252+
assert i.staggering == staggering
1253+
10961254
def test_dx2(self):
10971255
grid = Grid(shape=(4, 4))
10981256

0 commit comments

Comments
 (0)