From e93685dfbe1efda6d44bc56d40dd3ca055f7eed1 Mon Sep 17 00:00:00 2001 From: Fabio Luporini Date: Thu, 10 Sep 2026 08:39:18 +0100 Subject: [PATCH] compiler: Attach staggering metadata to IndexDerivative --- devito/finite_differences/differentiable.py | 35 +++- .../finite_differences/finite_difference.py | 3 +- devito/finite_differences/tools.py | 21 ++- tests/test_derivatives.py | 164 +++++++++++++++++- 4 files changed, 209 insertions(+), 14 deletions(-) diff --git a/devito/finite_differences/differentiable.py b/devito/finite_differences/differentiable.py index c38505414d..1a263c9945 100644 --- a/devito/finite_differences/differentiable.py +++ b/devito/finite_differences/differentiable.py @@ -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 @@ -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: @@ -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) @@ -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`. + """ + return self._staggering + @property def depth(self): iderivs = self.expr.find(IndexDerivative) @@ -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: diff --git a/devito/finite_differences/finite_difference.py b/devito/finite_differences/finite_difference.py index 73becb8b84..b9fcc26a12 100644 --- a/devito/finite_differences/finite_difference.py +++ b/devito/finite_differences/finite_difference.py @@ -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 = [] diff --git a/devito/finite_differences/tools.py b/devito/finite_differences/tools.py index 8c9304b126..8589fe4374 100644 --- a/devito/finite_differences/tools.py +++ b/devito/finite_differences/tools.py @@ -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: @@ -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 @@ -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): """ @@ -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): @@ -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 # Shift for side side = side or centered @@ -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): diff --git a/tests/test_derivatives.py b/tests/test_derivatives.py index 362d3d1201..89b2b25a41 100644 --- a/tests/test_derivatives.py +++ b/tests/test_derivatives.py @@ -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 ( @@ -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 @@ -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))