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+ )
1921from devito .finite_differences .tools import coeff_priority , make_shift_x0
2022from devito .logger import warning
2123from 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
0 commit comments