Skip to content

Commit e26adc2

Browse files
authored
Merge pull request #3001 from devitocodes/symmetric-interp-gradients
api: Symmetric interp gradients
2 parents da56fad + ec12fb1 commit e26adc2

6 files changed

Lines changed: 717 additions & 132 deletions

File tree

‎devito/finite_differences/differentiable.py‎

Lines changed: 73 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,9 @@
1515
# Moved in 1.13
1616
from sympy.core.basic import ordering_of_classes
1717

18-
from devito.finite_differences.interpolation import interp_at, post_x0_indices
18+
from devito.finite_differences.interpolation import (
19+
interp_at, interp_mapper, post_x0_indices
20+
)
1921
from devito.finite_differences.tools import coeff_priority, make_shift_x0
2022
from devito.logger import warning
2123
from devito.tools import (
@@ -487,16 +489,23 @@ def has_free(self, *patterns):
487489
return all(i in self.free_symbols for i in patterns)
488490

489491

490-
def highest_priority(diff_op):
491-
if not diff_op._args_diff:
492+
def highest_priority(diff_op, candidates=None):
493+
"""
494+
The Function whose location a product should be evaluated at.
495+
496+
`candidates` restricts the choice to a subset of the operands; without it
497+
the whole expression's differentiable arguments are considered.
498+
"""
499+
args_diff = diff_op._args_diff if candidates is None else tuple(candidates)
500+
if not args_diff:
492501
return diff_op
493502

494503
# We want to get the object with highest priority
495504
# We also need to make sure that the object with the largest
496505
# set of dimensions is used when multiple ones with the same
497506
# priority appear
498507
prio = lambda x: (getattr(x, '_fd_priority', 0), len(x.dimensions))
499-
prio_func = sorted(diff_op._args_diff, key=prio, reverse=True)[0]
508+
prio_func = sorted(args_diff, key=prio, reverse=True)[0]
500509

501510
# The highest priority must be a Function
502511
if not isinstance(prio_func, AbstractFunction):
@@ -664,6 +673,27 @@ def _gather_for_diff(self):
664673
other = self.func(*other)._eval_at(highest_priority(self))
665674
return self.func(other, *derivs)
666675

676+
@classmethod
677+
def _off_func(cls, a, func):
678+
"""
679+
Whether evaluating `a` at `func`'s location takes an interpolation.
680+
681+
A derivative is where its `x0` puts it, not where the field it
682+
differentiates lives: `div(v)` of a staggered velocity lands on the
683+
node, and reading its operand's staggering instead would re-associate a
684+
product that is already node-centred and interpolate it for nothing.
685+
A composite is off only if one of its own operands is, which is what
686+
makes a sum of derivatives -- a divergence, a trace -- read as the
687+
node-centred quantity it is.
688+
"""
689+
if isinstance(a, sympy.Derivative):
690+
source = post_x0_indices(a, func)
691+
elif isinstance(a, AbstractFunction) or not getattr(a, '_args_diff', ()):
692+
source = a.indices_ref
693+
else:
694+
return any(cls._off_func(i, func) for i in a._args_diff)
695+
return bool(interp_mapper(source, func.indices_ref, a.dimensions))
696+
667697
def _eval_at(self, func, interp_mode='direct', **kwargs):
668698
"""
669699
Evaluate a Mul at the location of `func`.
@@ -674,36 +704,52 @@ def _eval_at(self, func, interp_mode='direct', **kwargs):
674704
independently evaluated at `func`'s location via
675705
`Differentiable._eval_at`.
676706
677-
- `interp_mode='symmetric'`: when every Differentiable factor has a
678-
staggering different from `func`'s, apply the `I * (a * I^T * b)`
679-
form:
680-
681-
1. Pick a `block` location -- the highest-priority factor's
682-
staggering (NODE is the highest priority, so coefficient-like
683-
NODE factors win, as in the `I * C * I^T` elastic stiffness
684-
pattern). Each factor not at the block is brought there via
685-
`I^T` (an explicit 0-order FD interpolation operator).
686-
Derivatives additionally set `x0` on their own derivative
687-
dimensions to `func`'s indices.
707+
- `interp_mode='symmetric'`: the product is formed *away* from `func`
708+
and closed with a single interpolation, the `I * (a * I^T * b)` form:
709+
710+
1. Pick a `block` location -- the highest-priority staggering
711+
among the factors that are not already at `func`'s (NODE is the
712+
highest priority, so coefficient-like NODE factors win, as in
713+
the `I * C * I^T` elastic stiffness pattern). Each factor not at
714+
the block is brought there via `I^T` (an explicit 0-order FD
715+
interpolation operator). Derivatives additionally set `x0` on
716+
their own derivative dimensions to `func`'s indices.
688717
2. The product is formed at `block`'s location.
689718
3. The whole product is interpolated to `func` via `I` (an
690719
explicit 0-order FD operator).
691720
692-
When the trigger does not hold (e.g. some factor already matches
693-
`func`'s staggering), we fall back to `direct`.
721+
It takes *two* factors away from `func` to have something to
722+
re-associate. With one, the single interpolation `direct` puts on it
723+
is already the transpose-consistent form -- it is what makes the
724+
`i, j` entry of a stiffness matrix the transpose of its `j, i` entry
725+
-- so we fall back. With two, `direct` would interpolate each of them
726+
separately, and `I(a)*I(b)` is not `I(a*b)`: the discretized operator
727+
stops being the transpose of itself, which is invisible in a forward
728+
simulation and shows up as a first-order gradient in an adjoint one.
729+
730+
Which factors count is the other half of it. A derivative is where
731+
its `x0` puts it, not where its operand lives, so a product around a
732+
`div(v)` of a staggered velocity -- node-centred however its operands
733+
are staggered -- is left alone rather than re-associated onto the
734+
operands' location, which would replace a compact stencil with an
735+
interpolated one twice as wide.
694736
"""
695737
if interp_mode != 'symmetric':
696738
return super()._eval_at(func, **kwargs)
697739

698740
diff, other = split(self.args, lambda a: isinstance(a, Differentiable))
699741

700-
# Symmetric form requires every Differentiable factor to differ from
701-
# func; otherwise direct evaluation is cleaner and equivalent.
702-
if len(diff) < 2 or \
703-
any(a.staggered == func.staggered for a in diff):
704-
return super()._eval_at(func, **kwargs)
742+
# A single factor cannot be re-associated, and with everything already
743+
# on `func` there is no interpolation to place. The mode still has to
744+
# travel down: a product that cannot itself be re-associated is
745+
# routinely wrapped around one that can (a `dt` scaling, a sum of
746+
# per-component contractions), and dropping the mode here would silently
747+
# evaluate all of that in `direct`.
748+
off_func = [a for a in diff if self._off_func(a, func)]
749+
if len(off_func) < 2:
750+
return super()._eval_at(func, interp_mode=interp_mode, **kwargs)
705751

706-
block_indices = highest_priority(self).indices_ref
752+
block_indices = highest_priority(self, candidates=off_func).indices_ref
707753

708754
# Bring each factor to block's location (I^T where needed)
709755
new_factors = list(other)
@@ -1119,9 +1165,11 @@ def __new__(cls, *args, base=None, **kwargs):
11191165
except AttributeError:
11201166
# This might happen if e.g. one attempts a (re)construction with
11211167
# one sole argument. The (re)constructed EvalDerivative degenerates
1122-
# to an object of different type, in classic SymPy style. That's fine
1168+
# to an object of different type, in classic SymPy style. That's
1169+
# fine -- and a single argument that is itself a sum is the same
1170+
# story: a zero-order derivative whose weights collapse to one is
1171+
# the identity, so it comes back as the sum it was applied to.
11231172
assert len(args) <= 1
1124-
assert not obj.is_Add
11251173
return obj
11261174

11271175
return obj

‎devito/types/basic.py‎

Lines changed: 35 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
import abc
22
import inspect
33
import warnings
4-
from contextlib import suppress
4+
from contextlib import contextmanager, suppress
55
from ctypes import POINTER, Structure, _Pointer, c_char, c_char_p
66
from functools import cached_property, reduce
77
from operator import mul
@@ -1471,6 +1471,22 @@ def __getnewargs_ex__(self):
14711471
return args, kwargs
14721472

14731473

1474+
@contextmanager
1475+
def ignore_non_expr_deprecation():
1476+
"""
1477+
Suppress sympy's deprecation of non-`Expr` entries in a Matrix.
1478+
1479+
A Devito tensor is a Matrix of Devito objects, some of which are legitimately
1480+
not `Expr` (for example a serialization string), so the deprecation does not
1481+
apply to us. Sympy emits it from `_dod_to_DomainMatrix`, reached whenever a
1482+
tensor is built from components, and from `_unify_element_sympy`, reached on
1483+
scalar multiplication, so both are wrapped in this.
1484+
"""
1485+
with warnings.catch_warnings():
1486+
warnings.filterwarnings("ignore", category=SymPyDeprecationWarning)
1487+
yield
1488+
1489+
14741490
class AbstractTensor(sympy.ImmutableDenseMatrix, Basic, Pickable, Evaluable):
14751491

14761492
"""
@@ -1563,23 +1579,29 @@ def __subfunc_setup__(cls, *args, **kwargs):
15631579
return []
15641580

15651581
@classmethod
1566-
def _sympify(self, arg):
1582+
def _sympify(cls, arg):
15671583
# This is used internally by sympy to process arguments at rebuilt. And since
1568-
# some of our properties are non-sympyfiable we need to have a fallback
1584+
# some of our properties are non-sympyfiable we need to have a fallback.
1585+
# `strict` so that strings are left alone rather than parsed into Symbols,
1586+
# while plain numbers are turned into `Expr` as sympy expects (a Matrix
1587+
# holding non-`Expr` entries, such as a plain `int` 0, is deprecated)
15691588
try:
1570-
# Pure sympy object
1571-
return arg._sympy_()
1572-
except AttributeError:
1589+
return sympy.sympify(arg, strict=True)
1590+
except sympy.SympifyError:
15731591
return arg
15741592

15751593
@classmethod
1576-
def _eval_from_dok(cls, rows, cols, dok):
1577-
with warnings.catch_warnings():
1578-
warnings.filterwarnings(
1579-
"ignore",
1580-
category=SymPyDeprecationWarning
1581-
)
1582-
return super()._eval_from_dok(rows, cols, dok)
1594+
def _dod_to_DomainMatrix(cls, rows, cols, dod, types):
1595+
# Entry point of every matrix construction, from either a flat list
1596+
# (`_new`) or a dok (`_eval_from_dok`)
1597+
with ignore_non_expr_deprecation():
1598+
return super()._dod_to_DomainMatrix(rows, cols, dod, types)
1599+
1600+
@classmethod
1601+
def _unify_element_sympy(cls, rep, element):
1602+
# Entry point of scalar multiplication, e.g. `lam * tau`
1603+
with ignore_non_expr_deprecation():
1604+
return super()._unify_element_sympy(rep, element)
15831605

15841606
@property
15851607
def grid(self):

‎devito/types/equation.py‎

Lines changed: 21 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,15 @@ class Eq(sympy.Eq, Evaluable, Pickable):
3636
An ordered list of Dimensions that do not explicitly appear in either the
3737
left-hand side or in the right-hand side, but that should be honored when
3838
constructing an Operator.
39+
interp_mode : str, optional, default=None
40+
Overrides the Operator's `sym_opt={'interp-mode': ...}` for this
41+
equation only. An Operator generally wants one mode -- `'direct'`, which
42+
keeps finite-difference stencils compact -- while a single equation in
43+
it may need `'symmetric'`, which re-associates a product of staggered
44+
operands so that the discretized operator is the transpose of itself.
45+
A gradient accumulated alongside the propagation it is the adjoint of is
46+
the typical case: the stencils must stay compact, but the accumulation
47+
has to be an exact transpose or the gradient degrades to first order.
3948
4049
Examples
4150
--------
@@ -60,10 +69,10 @@ class Eq(sympy.Eq, Evaluable, Pickable):
6069
is_Reduction = False
6170

6271
__rargs__ = ('lhs', 'rhs')
63-
__rkwargs__ = ('subdomain', 'coefficients', 'implicit_dims')
72+
__rkwargs__ = ('subdomain', 'coefficients', 'implicit_dims', 'interp_mode')
6473

6574
def __new__(cls, lhs, rhs=0, subdomain=None, coefficients=None,
66-
implicit_dims=None, **kwargs):
75+
implicit_dims=None, interp_mode=None, **kwargs):
6776
if coefficients is not None:
6877
_ = deprecations.coeff_warn
6978
kwargs['evaluate'] = False
@@ -76,9 +85,15 @@ def __new__(cls, lhs, rhs=0, subdomain=None, coefficients=None,
7685
obj._subdomain = subdomain
7786
obj._substitutions = coefficients
7887
obj._implicit_dims = as_tuple(implicit_dims)
88+
obj._interp_mode = interp_mode
7989

8090
return obj
8191

92+
@property
93+
def interp_mode(self):
94+
"""Per-equation override of the Operator's `interp-mode`, or None."""
95+
return self._interp_mode
96+
8297
@classmethod
8398
def _apply_coeffs(cls, expr, coefficients):
8499
"""
@@ -108,14 +123,17 @@ def _evaluate(self, **kwargs):
108123
109124
The RHS of the Equation is evaluated at the indices of the LHS if required.
110125
"""
126+
if self._interp_mode is not None:
127+
kwargs['interp_mode'] = self._interp_mode
111128
try:
112129
lhs = self.lhs._evaluate(**kwargs)
113130
rhs = self.rhs._eval_at(self.lhs, **kwargs)._evaluate(**kwargs)
114131
except AttributeError:
115132
lhs, rhs = self._evaluate_args(**kwargs)
116133
eq = self.func(lhs, rhs, subdomain=self.subdomain,
117134
coefficients=self.substitutions,
118-
implicit_dims=self._implicit_dims)
135+
implicit_dims=self._implicit_dims,
136+
interp_mode=self._interp_mode)
119137

120138
return eq
121139

0 commit comments

Comments
 (0)