Skip to content
Closed
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
35 changes: 29 additions & 6 deletions devito/finite_differences/differentiable.py
Original file line number Diff line number Diff line change
Expand Up @@ -1048,9 +1048,9 @@ def value(self, idx):
class IndexDerivative(IndexSum):

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

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

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

obj._deriv_order = deriv_order
obj._staggering = staggering

return obj

@cached_property
def _metadata(self):
# SymPy's canonical sorting also compares the hashable content directly.
# Use comparable objects, including empty tuples for unknown metadata
return (sympy.Dict(*self.mapper.items()),
sympy.Tuple(*as_tuple(self.deriv_order)),
sympy.Tuple(*as_tuple(self.staggering)))

def _hashable_content(self):
return super()._hashable_content() + (self.mapper,)
return super()._hashable_content() + self._metadata

def compare(self, other):
if self is other:
Expand All @@ -1086,7 +1095,8 @@ def compare(self, other):
n2 = other.__class__
if n1.__name__ == n2.__name__:
return (self.weights.compare(other.weights) or
self.base.compare(other.base))
self.base.compare(other.base) or
super().compare(other))
else:
return super().compare(other)

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

@property
def staggering(self):
"""
The requested evaluation staggering relative to the input lattice:
`centered` on that lattice, `staggered` halfway between its points,
or None for other or unknown evaluation locations.

This classification is independent of differential order, stencil bias,
and transposition. In particular, `centered` does not imply symmetric
weights, and interpolation is identified separately by `deriv_order == 0`.
"""
Comment on lines +1125 to +1133

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.

Nitpick: excessively verbose docstring.

return self._staggering

@property
def depth(self):
iderivs = self.expr.find(IndexDerivative)
Expand Down Expand Up @@ -1288,8 +1311,8 @@ def _diff2sympy(obj):

# Handle special objects
if isinstance(obj, DiffDerivative):
return IndexDerivative(*args, obj.mapper,
deriv_order=obj.deriv_order), True
kwargs = {i: getattr(obj, i) for i in obj.__rkwargs__}
return IndexDerivative(*args, obj.mapper, **kwargs), True

# Handle generic objects such as arithmetic operations
try:
Expand Down
3 changes: 2 additions & 1 deletion devito/finite_differences/finite_difference.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,7 +236,8 @@ def make_derivative(expr, dim, fd_order, deriv_order, side, matvec, x0, coeffici
expr = expr._evaluate(expand=False)

deriv = DiffDerivative(
expr*weights, {dim: indices.free_dim}, deriv_order=deriv_order
expr*weights, {dim: indices.free_dim}, deriv_order=deriv_order,
staggering=indices.staggering
)
else:
terms = []
Expand Down
21 changes: 17 additions & 4 deletions devito/finite_differences/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,9 +146,13 @@ class IndexSet(tuple):

"""
The points of a finite-difference expansion.

`staggering` records the scheme's requested evaluation staggering relative
to the input lattice: `centered`, `staggered`, or None if unknown. Index
coordinate changes preserve this classification.
"""

def __new__(cls, dim, indices=None, expr=None, fd=None):
def __new__(cls, dim, indices=None, expr=None, fd=None, staggering=None):
assert indices is not None or expr is not None

if fd is None:
Expand All @@ -167,6 +171,7 @@ def __new__(cls, dim, indices=None, expr=None, fd=None):
obj.dim = dim
obj.expr = expr
obj.free_dim = fd
obj.staggering = staggering

return obj

Expand Down Expand Up @@ -203,7 +208,8 @@ def transpose(self):
except AttributeError:
expr = None

return IndexSet(self.dim, indices, expr=expr, fd=free_dim)
return IndexSet(self.dim, indices, expr=expr, fd=free_dim,
staggering=self.staggering)

def shift(self, v):
"""
Expand All @@ -216,7 +222,8 @@ def shift(self, v):
except TypeError:
expr = None

return IndexSet(self.dim, indices, expr=expr, fd=self.free_dim)
return IndexSet(self.dim, indices, expr=expr, fd=self.free_dim,
staggering=self.staggering)


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

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

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.

This seems oddly convoluted? Can this really not be simplified at all?


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

return IndexSet(dim, expr=iexpr), x0
return IndexSet(dim, expr=iexpr, staggering=staggering), x0


def make_shift_x0(shift, ndim):
Expand Down
164 changes: 161 additions & 3 deletions tests/test_derivatives.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import numpy as np
import pytest
from sympy import Float, Symbol, diff, simplify, sympify
from sympy import Float, S, Symbol, diff, simplify, sympify

from conftest import assert_structure
from devito import (
Expand All @@ -10,9 +10,12 @@
)
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, IndexSum, Weights, diff2sympy,
interp_for_fd
)
from devito.symbolics import indexify, retrieve_indexed
from devito.finite_differences.tools import generate_indices
from devito.ir.equations.algorithms import lower_exprs
from devito.symbolics import indexify, retrieve_indexed, search, uxreplace
from devito.types.dimension import StencilDimension
from devito.warnings import DevitoWarning

Expand Down Expand Up @@ -1093,6 +1096,161 @@ def test_index_derivative(self):

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

@pytest.mark.parametrize('staggered', [None, x, y])
@pytest.mark.parametrize('deriv_order,fd_order', [
(0, 4), (0, 16), (1, 2), (1, 4), (1, 16), (2, 4), (2, 16)
])
@pytest.mark.parametrize('offset,staggering', [
(0, 'centered'), (S.Half, 'staggered'), (-S.Half, 'staggered'),
(1, 'centered'), (S.One/4, None)
])
def test_index_derivative_staggering(self, staggered, deriv_order, fd_order,
offset, staggering):
grid = Grid(shape=(10, 10))
x, y = grid.dimensions
staggered = staggered(grid) if staggered else NODE
f = Function(name='f', grid=grid, space_order=16, staggered=staggered)
x0 = {x: f.indices_ref[x] + offset*x.spacing}
deriv = f.diff(x, deriv_order=deriv_order, fd_order=fd_order, x0=x0)

evaluated = deriv._evaluate(expand=False)
if deriv_order == 0 and offset == 0:
assert evaluated == f
return

lowered = lower_exprs(diff2sympy(evaluated))
for i in (evaluated, lowered):
assert isinstance(i, IndexDerivative)
assert i.staggering == staggering
assert i.deriv_order == deriv_order

# Classifying the staggering must not change the discrete stencil
sd, = evaluated.dimensions
terms = [w*evaluated.base.subs(sd, i) for w, i in
zip(evaluated.weights.function.weights, sd.range, strict=True)]
assert simplify(sum(terms) - deriv.evaluate) == 0

@pytest.mark.parametrize('offsets,staggerings,fd_order,weights', [
((0, S.Half), ('centered', 'staggered'), 2, None),
((S.Half, S.One/4), ('staggered', None), 4,
[S.One/24, -S(9)/8, S(9)/8, -S.One/24])
])
@pytest.mark.parametrize('lowered', [False, True])
def test_index_derivative_staggering_identity(self, offsets, staggerings, fd_order,
weights, lowered):
grid = Grid(shape=(10,))
x, = grid.dimensions
f = Function(name='f', grid=grid, space_order=4)
derivs = [f.dx(fd_order=fd_order, x0={x: x + i*x.spacing}, weights=weights)
for i in offsets]
derivs = [i._evaluate(expand=False) for i in derivs]
if lowered:
derivs = [lower_exprs(diff2sympy(i)) for i in derivs]
a, b = derivs

# Distinct staggerings can generate exactly the same discrete stencil
assert tuple(i.staggering for i in derivs) == staggerings
assert a.args == b.args
assert a.mapper == b.mapper
assert a.weights.function.weights == b.weights.function.weights
assert a != b
assert len({a, b}) == 2
assert a.compare(b) == -b.compare(a) != 0
assert set((a + b).args) == {a, b}
assert a.evaluate == b.evaluate

@pytest.mark.parametrize('offset,staggering', [
(-S.Half, 'staggered'), (0, 'centered')
])
@pytest.mark.parametrize('side', [None, centered, left, right])
@pytest.mark.parametrize('transpose', [False, True])
def test_index_derivative_staggering_transpose(self, offset, staggering, side,
transpose):
grid = Grid(shape=(10,))
x, = grid.dimensions
f = Function(name='f', grid=grid, space_order=4, staggered=x)
deriv = f.dx(x0={x: f.indices_ref[x] + offset*x.spacing}, side=side)
if transpose:
deriv = deriv.T
evaluated = deriv._evaluate(expand=False)
assert evaluated.staggering == staggering
assert lower_exprs(diff2sympy(evaluated)).staggering == staggering
assert simplify(evaluated.evaluate - deriv.evaluate) == 0

@pytest.mark.parametrize('dim', [x, y])
@pytest.mark.parametrize('offset,staggering', [
(-S.Half, 'staggered'), (0, 'centered')
])
def test_index_derivative_staggering_nested(self, dim, offset, staggering):
grid = Grid(shape=(10, 10))
x, y = grid.dimensions
d = dim(grid)
f = Function(name='f', grid=grid, space_order=4)
deriv = f.dx(x0={x: x + x.spacing/2})
deriv = deriv.diff(d, deriv_order=2, x0={d: d + offset*d.spacing})

evaluated = deriv._evaluate(expand=False)
lowered = lower_exprs(diff2sympy(evaluated))
for expr in (evaluated, lowered):
iderivs = search(expr, IndexDerivative)
assert len(iderivs) == 2
assert {(i.deriv_order, i.staggering) for i in iderivs} == {
(1, 'staggered'), (2, staggering)
}

@pytest.mark.parametrize('lowered', [False, True])
def test_index_derivative_staggering_rebuild(self, lowered):
grid = Grid(shape=(10,))
x, = grid.dimensions
f = Function(name='f', grid=grid, space_order=4)
g = Function(name='g', grid=grid, space_order=4)
expr = f.dx(x0={x: x + x.spacing/2})._evaluate(expand=False)
sd, = expr.dimensions
base = g.subs(x, x + sd*x.spacing)
if lowered:
expr = lower_exprs(diff2sympy(expr))
base = lower_exprs(base)

for rebuilt in (expr.func(base*expr.weights), expr.subs(expr.base, base),
expr.xreplace({expr.base: base}),
uxreplace(expr, {expr.base: base})):
assert rebuilt.base == base
assert rebuilt.staggering == 'staggered'
assert rebuilt.deriv_order == 1

# Missing metadata must stay distinct from a known centered interpolation
unknown = expr._rebuild(staggering=None, deriv_order=None)
on_grid = expr._rebuild(staggering='centered', deriv_order=0)
assert unknown.staggering is None
assert len({expr, unknown, on_grid}) == 3
assert unknown.compare(on_grid) == -on_grid.compare(unknown) != 0
assert set((expr + unknown + on_grid).args) == {expr, unknown, on_grid}

def test_index_derivative_staggering_unknown(self):
grid = Grid(shape=(10, 10))
x, y = grid.dimensions
f = Function(name='f', grid=grid, space_order=4)
exprs = (f.dx45, f.dx(x0={x: 1}))
for expr in exprs:
lowered = lower_exprs(diff2sympy(expr._evaluate(expand=False)))
iderivs = search(lowered, IndexDerivative)
assert iderivs
assert all(i.staggering is None for i in iderivs)

@pytest.mark.parametrize('offset,staggering', [
(0, 'centered'), (1.0, 'centered'), (-1.0, 'centered'),
(0.5, 'staggered'), (-1.5, 'staggered'), (0.25, None), (0.50000001, None)
])
def test_index_set_staggering(self, offset, staggering):
grid = Grid(shape=(10,))
x, = grid.dimensions
f = Function(name='f', grid=grid, space_order=4, staggered=x)
indices, _ = generate_indices(f, x, 4,
x0={x: f.indices_ref[x] + offset*x.spacing})
for i in (indices, indices.transpose(), indices.shift(-x.spacing/2),
indices.transpose().shift(-x.spacing/2)):
assert i.staggering == staggering

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

Expand Down
Loading