Skip to content

Commit e96d4ca

Browse files
committed
dsl: Differentiate a mixed-staggering sum term by term
`Add` reports its first argument's `indices_ref`, so a sum whose terms sit at different staggered locations names a position only one of them has, and `x0` gets resolved against it for all of them. The shear strain `v_x.dy + v_y.dx` of a staggered velocity is the canonical case: both terms land on the cell corner, so a shift onto it should be a no-op, and instead each picked up a spurious one. Differentiation is linear at every order, so split such a sum in `Derivative._eval_fd`. Relative error on `D(a+b)` against `D(a) + D(b)` was 0.63 at order 0, 1.20 at order 1 and 0.95 at order 2, with `expand=False` at order 2 returning exactly zero. `generic_derivative` also short-circuited a zeroth order derivative only when `x0` was empty, building a stencil around an expression already sitting at `x0`. `index_at` answers where an expression sits, and both call sites use it.
1 parent 8f1def8 commit e96d4ca

3 files changed

Lines changed: 64 additions & 3 deletions

File tree

devito/finite_differences/derivative.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414
from devito.warnings import warn
1515

1616
from .differentiable import Add, Differentiable, Mul, diffify, interp_for_fd
17-
from .finite_difference import cross_derivative, generic_derivative
17+
from .finite_difference import cross_derivative, generic_derivative, index_at
1818
from .rsfd import d45
1919
from .tools import direct, transpose
2020

@@ -575,6 +575,12 @@ def _eval_fd(self, expr, **kwargs):
575575
shited derivative.
576576
- 4: Apply substitutions.
577577
"""
578+
# Differentiation is linear, and a sum of terms at different staggered
579+
# locations must use it: `Add` reports its first argument's location,
580+
# so `x0` would shift the other terms off the point they sat at.
581+
if expr.is_Add and any(index_at(expr, d) is None for d in self.dims):
582+
return expr.func(*[self._eval_fd(a, **kwargs) for a in expr.args])
583+
578584
# Step 1: Evaluate non-derivative x0. We currently enforce a simple 2nd order
579585
# interpolation to avoid very expensive finite differences on top of it
580586
x0_deriv = self._filter_dims(self.x0)

devito/finite_differences/finite_difference.py

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,18 @@ def cross_derivative(expr, dims, fd_order, deriv_order, x0=None, side=None, **kw
100100
return expr
101101

102102

103+
def index_at(expr, dim):
104+
"""
105+
Where `expr` sits along `dim`, or None if its terms disagree.
106+
"""
107+
try:
108+
indices = {i.indices_ref[dim] for i in expr.args} if expr.is_Add \
109+
else {expr.indices_ref[dim]}
110+
except (AttributeError, KeyError, IndexError, TypeError):
111+
return None
112+
return indices.pop() if len(indices) == 1 else None
113+
114+
103115
@check_input
104116
def generic_derivative(expr, dim, fd_order, deriv_order, matvec=direct, x0=None,
105117
coefficients='taylor', expand=True, weights=None, side=None):
@@ -139,8 +151,9 @@ def generic_derivative(expr, dim, fd_order, deriv_order, matvec=direct, x0=None,
139151
if deriv_order == 1 and fd_order == 2 and side is None:
140152
fd_order = 1
141153

142-
# Zeroth order derivative is just the expression itself if not shifted
143-
if deriv_order == 0 and not x0:
154+
# Zeroth order is the identity when `expr` already sits at `x0`, not a
155+
# stencil centred there.
156+
if deriv_order == 0 and (not x0 or index_at(expr, dim) == x0.get(dim)):
144157
return expr
145158

146159
# Enforce stable time coefficients

tests/test_derivatives.py

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1461,3 +1461,45 @@ def test_unevaluated(self):
14611461
assert Derivative(self.x, self.t)
14621462
assert Derivative(self.x, self.y, self.t)
14631463
assert Derivative(self.x, (self.x, 0))
1464+
1465+
1466+
@pytest.mark.parametrize('expand', [True, False])
1467+
@pytest.mark.parametrize('deriv_order', [0, 1, 2])
1468+
def test_deriv_sum_mixed_staggering(expand, deriv_order):
1469+
"""
1470+
A shifted derivative is linear: `D(a + b) == D(a) + D(b)`, at every order.
1471+
1472+
It broke when the terms sat at different staggered locations, `Add`
1473+
reporting only its first argument's: the shear strain `vx.dy + vy.dx` of a
1474+
staggered velocity read 0.63 relative error at order 0, 1.20 at order 1 and
1475+
0.95 at order 2, `expand=False` at order 2 returning exactly zero.
1476+
"""
1477+
so = 8
1478+
grid = Grid(shape=(41, 41), extent=(40., 40.))
1479+
x, y = grid.dimensions
1480+
1481+
vx = Function(name='vx', grid=grid, space_order=so, staggered=x)
1482+
vy = Function(name='vy', grid=grid, space_order=so, staggered=y)
1483+
out = Function(name='out', grid=grid, space_order=so, staggered=(x, y))
1484+
1485+
rng = np.random.default_rng(3)
1486+
for f in (vx, vy):
1487+
f.data[:] = rng.normal(size=f.shape)
1488+
1489+
def shifted(expr):
1490+
return expr.diff(y, deriv_order=deriv_order, fd_order=2,
1491+
x0={y: y + y.spacing/2})
1492+
1493+
def run(expr):
1494+
out.data[:] = 0.
1495+
Operator(Eq(out, expr), opt=('advanced', {'expand': expand})).apply()
1496+
return np.array(out.data)
1497+
1498+
s = slice(so + 3, -(so + 3))
1499+
together = run(shifted(vx.dy + vy.dx))[s, s]
1500+
apart = (run(shifted(vx.dy)) + run(shifted(vy.dx)))[s, s]
1501+
1502+
assert np.linalg.norm(apart) > 0
1503+
# float32 reassociation only: the two forms sum the same terms in a
1504+
# different order
1505+
assert np.linalg.norm(together - apart) / np.linalg.norm(apart) < 1e-5

0 commit comments

Comments
 (0)