Skip to content

Commit 6573324

Browse files
committed
api: Add Derivative halo=0 for SubDomain-restricted derivatives and adjoints
`expr.dx(halo=0)` treats `expr` as zero outside the SubDomain of its equation, and `.T` keeps the flag. Hence `Inc(q_bar, out_bar.dx(halo=0).T, subdomain=S)` is the adjoint of `Eq(out, q.dx, subdomain=S)`, i.e. D^T R^T with R the restriction to S, without a user-side zero-padded work field. Equations without `halo=0` derivatives, or without a SubDomain, are evaluated as before. - An equation evaluates its rhs with `_eval_at(lhs, subdomain=S)`. A `halo=0` derivative records S, and multiplies each stencil tap by a branch-free integer MIN/MAX indicator of the grid point it reads being in S (SubDomain.indicator), which invariant hoisting computes once. Along the Dimensions it does not differentiate, it restricts its argument to S at the evaluation point. In a sum with such derivatives, the other terms are restricted to S at the evaluation point (SubDomain.restrict); factors are not. - The evaluated derivative records its stencil radius (halo_radius), and the equation iterates over S grown by the largest one (SubDomain.grow, a SubDomain with a parent). Tensor equations aggregate their components. - It works for expanded, unexpanded and staggered derivatives. Derivatives of `halo=0` derivatives, MultiSubDomains and methods other than FD raise NotImplementedError. Tests: dot tests of `Eq(out, g.dx, subdomain=S)` against `Inc(g_bar, out_bar.dx(halo=0).T, subdomain=S)` for left/right/middle SubDomains, first and second derivatives, orders 2/4/8, Eq and Inc, 2D and unexpanded forms, in 1D and 3D; staggered dot tests against `-out_bar.dx(halo=0)`, as in elastic adjoints, including a staggered vector equation; sums, factors and vector equations; SubDomain growth; a forward check; and a one-face CPML operator.
1 parent c5d3325 commit 6573324

8 files changed

Lines changed: 669 additions & 39 deletions

File tree

‎devito/finite_differences/derivative.py‎

Lines changed: 93 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -102,7 +102,7 @@ def _fd_priority(self):
102102

103103
__rargs__ = ('expr', '*dims')
104104
__rkwargs__ = ('side', 'deriv_order', 'fd_order', 'transpose', '_ppsubs',
105-
'x0', 'method', 'weights')
105+
'x0', 'method', 'weights', 'halo', 'subdomain')
106106

107107
def __new__(cls, expr, *dims, **kwargs):
108108
# Validate the input arguments `expr`, `dims` and `deriv_order`
@@ -161,6 +161,8 @@ def __new__(cls, expr, *dims, **kwargs):
161161
obj._transpose = kwargs.get("transpose", direct)
162162
obj._method = kwargs.get("method", 'FD')
163163
obj._weights = cls._process_weights(**kwargs)
164+
obj._halo = cls._validate_halo(kwargs.get("halo"))
165+
obj._subdomain = kwargs.get("subdomain")
164166

165167
ppsubs = kwargs.get("subs", kwargs.get("_ppsubs", []))
166168
processed = []
@@ -177,6 +179,16 @@ def __new__(cls, expr, *dims, **kwargs):
177179

178180
return obj
179181

182+
@staticmethod
183+
def _validate_halo(halo):
184+
"""
185+
Validate `halo`. Only None (read the argument everywhere) and 0 (treat
186+
the argument as zero outside the equation's SubDomain) are supported.
187+
"""
188+
if halo not in (None, 0):
189+
raise ValueError(f"Expected halo=None or halo=0, got halo={halo}")
190+
return halo
191+
180192
@staticmethod
181193
def _validate_expr(expr):
182194
"""
@@ -325,6 +337,8 @@ def _process_weights(cls, **kwargs):
325337
def __call__(self, x0=None, fd_order=None, side=None, method=None, **kwargs):
326338
weights = kwargs.get('weights', kwargs.get('w'))
327339
rkw = {}
340+
if 'halo' in kwargs:
341+
rkw['halo'] = kwargs['halo']
328342
if side is not None:
329343
rkw['side'] = side
330344
if method is not None:
@@ -457,6 +471,26 @@ def side(self):
457471
def transpose(self):
458472
return self._transpose
459473

474+
@property
475+
def halo(self):
476+
"""
477+
None if the argument is read everywhere, 0 if it is treated as zero
478+
outside the SubDomain of the equation the Derivative belongs to.
479+
"""
480+
return self._halo
481+
482+
@property
483+
def subdomain(self):
484+
"""
485+
With halo=0, the SubDomain outside of which the argument is treated as
486+
zero, set upon evaluation within an equation restricted to it.
487+
"""
488+
return self._subdomain
489+
490+
@cached_property
491+
def _has_zero_halo(self):
492+
return self.halo is not None or self.expr._has_zero_halo
493+
460494
@property
461495
def is_TimeDependent(self):
462496
return self.expr.is_TimeDependent
@@ -481,26 +515,21 @@ def T(self):
481515

482516
return self._rebuild(transpose=adjoint)
483517

484-
def _eval_at(self, func, interp_mode='direct', **kwargs):
518+
def _eval_at(self, func, interp_mode='direct', subdomain=None, **kwargs):
485519
"""
486520
Evaluates the derivative at the location of `func`. It is necessary for staggered
487521
setup where one could have Eq(u(x + h_x/2), v(x).dx)) in which case v(x).dx
488522
has to be computed at x=x + h_x/2.
523+
524+
With halo=0, the argument is treated as zero outside `subdomain`, which
525+
the Derivative records.
489526
"""
490-
# No staggering, don't waste time
491-
if not self.expr.staggered and not func.staggered:
492-
return self
493-
# If an x0 already exists or evaluating at the same function (i.e u = u.dx)
494-
# do not overwrite it
495-
if self.x0 or self.side is not None or func.function is self.expr.function:
496-
return self
497-
# For basic equation of the form f = Derivative(g, ...) we can just
498-
# compare staggering
499-
if self.expr.staggered == func.staggered and self.expr.is_Function:
500-
return self
501-
# Time derivatives are not affected by space staggering
502-
if all(d.is_Time for d in self.dims):
503-
return self
527+
rkw = {}
528+
if subdomain is not None and self.halo is not None:
529+
rkw['subdomain'] = subdomain
530+
531+
if not self._is_relocated(func):
532+
return self._rebuild(**rkw) if rkw else self
504533

505534
# Check if x0's keys come from a DerivedDimension
506535
x0 = func.indices_ref.getters
@@ -519,7 +548,7 @@ def _eval_at(self, func, interp_mode='direct', **kwargs):
519548
# e.g f.dx(x0={x: x + h_x/2}).subs({x: ix})
520549
psubs[sd] = d
521550
nx0[sd] = nx0.pop(d)._subs(d, sd)
522-
rkw = {'x0': nx0}
551+
rkw['x0'] = nx0
523552
if psubs:
524553
rkw['subs'] = (psubs,)
525554

@@ -534,7 +563,8 @@ def _eval_at(self, func, interp_mode='direct', **kwargs):
534563
return self._rebuild(self.expr, **rkw)
535564
args = [self.expr.func(*v) for v in mapper.values()]
536565
args.extend([a for a in self.expr.args if a not in self.expr._args_diff])
537-
args = [self._rebuild(a)._eval_at(func, interp_mode=interp_mode, **kwargs)
566+
args = [self._rebuild(a)._eval_at(func, interp_mode=interp_mode,
567+
subdomain=subdomain, **kwargs)
538568
for a in args]
539569
return self.expr.func(*args)
540570
elif self.expr.is_Mul:
@@ -549,6 +579,25 @@ def _eval_at(self, func, interp_mode='direct', **kwargs):
549579
# the expression as is.
550580
return self._rebuild(self.expr, **rkw)
551581

582+
def _is_relocated(self, func):
583+
"""
584+
True if the Derivative must be evaluated at the location of `func`, e.g.
585+
with `func` and the argument staggered apart.
586+
"""
587+
# No staggering, don't waste time
588+
if not self.expr.staggered and not func.staggered:
589+
return False
590+
# If an x0 already exists or evaluating at the same function (i.e u = u.dx)
591+
# do not overwrite it
592+
if self.x0 or self.side is not None or func.function is self.expr.function:
593+
return False
594+
# For basic equation of the form f = Derivative(g, ...) we can just
595+
# compare staggering
596+
if self.expr.staggered == func.staggered and self.expr.is_Function:
597+
return False
598+
# Time derivatives are not affected by space staggering
599+
return not all(d.is_Time for d in self.dims)
600+
552601
def _evaluate(self, **kwargs):
553602
# Evaluate finite-difference.
554603
# NOTE: `evaluate` and `_eval_fd` split for potential future different
@@ -581,6 +630,13 @@ def _eval_fd(self, expr, **kwargs):
581630
if expr.is_Add and any(len(indices_at(expr, d)) > 1 for d in self.dims):
582631
return expr.func(*[self._eval_fd(a, **kwargs) for a in expr.args])
583632

633+
# The SubDomain mask read by a halo=0 derivative can't be shifted again
634+
if any(d.halo is not None for d in expr.find(Derivative)):
635+
raise NotImplementedError(
636+
f"{self} differentiates a derivative with halo=0, which is not "
637+
"supported"
638+
)
639+
584640
# Step 1: Evaluate non-derivative x0. We currently enforce a simple 2nd order
585641
# interpolation to avoid very expensive finite differences on top of it
586642
x0_deriv = self._filter_dims(self.x0)
@@ -598,6 +654,15 @@ def _eval_fd(self, expr, **kwargs):
598654
# otherwise an IndexSum will returned
599655
expand = kwargs.get('expand', True)
600656

657+
# With halo=0, `expr` is treated as zero outside the equation's SubDomain
658+
subdomain = self.subdomain
659+
if subdomain is not None and (subdomain.is_MultiSubDomain or
660+
self.method != 'FD'):
661+
raise NotImplementedError(
662+
f"halo=0 is only supported with method='FD' on a SubDomain, not "
663+
f"with method={self.method} on {subdomain}"
664+
)
665+
601666
# Step 3: Evaluate FD of the new expression
602667
if self.method == 'RSFD':
603668
assert len(self.dims) == 1
@@ -607,18 +672,26 @@ def _eval_fd(self, expr, **kwargs):
607672
assert self.method == 'FD'
608673
res = cross_derivative(expr, self.dims, self.fd_order, self.deriv_order,
609674
matvec=self.transpose, x0=x0_deriv, expand=expand,
610-
side=self.side, weights=self.weights)
675+
side=self.side, weights=self.weights,
676+
subdomain=subdomain)
611677
else:
612678
assert self.method == 'FD'
613679
res = generic_derivative(expr, self.dims[0], self.fd_order[0],
614680
self.deriv_order[0], weights=self.weights,
615681
side=self.side, matvec=self.transpose,
616-
x0=self.x0, expand=expand)
682+
x0=self.x0, expand=expand,
683+
subdomain=subdomain)
617684

618685
# Step 4: Apply substitutions
619686
for e in self._ppsubs:
620687
res = res.xreplace(e)
621688

689+
# With `halo=0`, along the Dimensions it does not differentiate, the
690+
# argument is read at the evaluation point, which lies outside the
691+
# SubDomain wherever other derivatives extend the equation along them
692+
if subdomain is not None:
693+
res = subdomain.restrict(res, exclude={d.root for d in self.dims})
694+
622695
return res
623696

624697
def _eval_expand_nest(self, **hints):

‎devito/finite_differences/differentiable.py‎

Lines changed: 55 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -191,6 +191,24 @@ def _eval_at(self, func, **kwargs):
191191
for a in self.args # false positive: lambda is invoked in-place
192192
])
193193

194+
@cached_property
195+
def _has_zero_halo(self):
196+
"""True if the expression has derivatives with `halo=0`."""
197+
return any(a._has_zero_halo for a in self._args_diff)
198+
199+
@cached_property
200+
def halo_radius(self):
201+
"""
202+
The largest stencil radius, per root Dimension, of the evaluated
203+
derivatives with `halo=0` in the expression. The expression is nonzero up
204+
to that many points past the SubDomain their argument is restricted to.
205+
206+
For example, with `S` a SubDomain restricting `x` and an 8th-order `g`,
207+
`g.dx(halo=0)` evaluated in an equation on `S` has a halo radius of
208+
`{x: 4}`: the equation must iterate over `S` extended by 4 points.
209+
"""
210+
return merge_halo_radius(a.halo_radius for a in self._args_diff)
211+
194212
def _subs(self, old, new, **hints):
195213
if old == self:
196214
return new
@@ -540,6 +558,18 @@ def highest_priority(diff_op, candidates=None):
540558
return prio_func
541559

542560

561+
def merge_halo_radius(radii):
562+
"""
563+
Merge the halo radii `radii`, each a mapping from root Dimension to radius,
564+
keeping the largest radius per Dimension.
565+
"""
566+
radius = {}
567+
for i in radii:
568+
for d, r in i.items():
569+
radius[d] = max(radius.get(d, 0), r)
570+
return frozendict(radius)
571+
572+
543573
class DifferentiableOp(Differentiable):
544574

545575
__sympy_class__ = None
@@ -640,6 +670,19 @@ def __new__(cls, *args, **kwargs):
640670

641671
return super().__new__(cls, *args, **kwargs)
642672

673+
def _eval_at(self, func, subdomain=None, **kwargs):
674+
"""
675+
Evaluate the sum at the location of `func`.
676+
677+
The derivatives with `halo=0` extend the sum past `subdomain`, so the
678+
other terms are restricted to `subdomain` at the evaluation point.
679+
"""
680+
expr = super()._eval_at(func, subdomain=subdomain, **kwargs)
681+
if subdomain is None or not expr.is_Add or not expr._has_zero_halo:
682+
return expr
683+
halo = {a for a in expr._args_diff if a._has_zero_halo}
684+
return self.func(*[a if a in halo else subdomain.restrict(a) for a in expr.args])
685+
643686

644687
class Mul(DifferentiableOp, sympy.Mul):
645688
__sympy_class__ = sympy.Mul
@@ -1276,6 +1319,14 @@ def _subs(self, old, new, **hints):
12761319

12771320
class DiffDerivative(IndexDerivative, DifferentiableOp):
12781321

1322+
__rkwargs__ = IndexDerivative.__rkwargs__ + ('halo_radius',)
1323+
1324+
def __new__(cls, *args, halo_radius=None, **kwargs):
1325+
obj = super().__new__(cls, *args, **kwargs)
1326+
# With `halo=0`, the stencil radius (see `Differentiable.halo_radius`)
1327+
obj.halo_radius = frozendict(halo_radius or {})
1328+
return obj
1329+
12791330
def _eval_at(self, func, **kwargs):
12801331
# Like EvalDerivative, a DiffDerivative must have already been evaluated
12811332
# at a valid x0 and should not be re-evaluated at a different location
@@ -1291,9 +1342,9 @@ class EvalDerivative(DifferentiableOp, sympy.Add):
12911342

12921343
is_commutative = True
12931344

1294-
__rkwargs__ = ('base',)
1345+
__rkwargs__ = ('base', 'halo_radius')
12951346

1296-
def __new__(cls, *args, base=None, **kwargs):
1347+
def __new__(cls, *args, base=None, halo_radius=None, **kwargs):
12971348
kwargs['evaluate'] = False
12981349

12991350
# a+0 -> a
@@ -1309,6 +1360,8 @@ def __new__(cls, *args, base=None, **kwargs):
13091360
# In some rare cases (rebuild?) base may be obj itself
13101361
base = base.base
13111362
obj.base = base
1363+
# With `halo=0`, the stencil radius (see `Differentiable.halo_radius`)
1364+
obj.halo_radius = frozendict(halo_radius or {})
13121365
except AttributeError:
13131366
# This might happen if e.g. one attempts a (re)construction with
13141367
# one sole argument. The (re)constructed EvalDerivative degenerates

‎devito/finite_differences/finite_difference.py‎

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -125,7 +125,8 @@ def index_at(expr, dim):
125125

126126
@check_input
127127
def generic_derivative(expr, dim, fd_order, deriv_order, matvec=direct, x0=None,
128-
coefficients='taylor', expand=True, weights=None, side=None):
128+
coefficients='taylor', expand=True, weights=None, side=None,
129+
subdomain=None):
129130
"""
130131
Arbitrary-order derivative of a given expression.
131132
@@ -151,6 +152,10 @@ def generic_derivative(expr, dim, fd_order, deriv_order, matvec=direct, x0=None,
151152
expand : bool, optional, default=True
152153
If True, the derivative is fully expanded as a sum of products,
153154
otherwise an IndexSum is returned.
155+
subdomain : SubDomain, optional
156+
If given, `expr` is treated as zero outside `subdomain` (see
157+
`Derivative`'s `halo`): each stencil tap is multiplied by the
158+
indicator of `subdomain` at the point it reads.
154159
155160
Returns
156161
-------
@@ -171,7 +176,7 @@ def generic_derivative(expr, dim, fd_order, deriv_order, matvec=direct, x0=None,
171176
coefficients = 'taylor' if dim.is_Time else expr.coefficients
172177

173178
return make_derivative(expr, dim, fd_order, deriv_order, side,
174-
matvec, x0, coefficients, expand, weights)
179+
matvec, x0, coefficients, expand, weights, subdomain)
175180

176181

177182
# Backward compatibility
@@ -180,7 +185,7 @@ def first_derivative(expr, dim, fd_order, **kwargs):
180185

181186

182187
def make_derivative(expr, dim, fd_order, deriv_order, side, matvec, x0, coefficients,
183-
expand, weights=None):
188+
expand, weights=None, subdomain=None):
184189
# Always expand time derivatives to avoid issue with buffering and streaming.
185190
# Time derivative are almost always short stencils and won't benefit from
186191
# unexpansion in the rare case the derivative is not evaluated for time stepping.
@@ -222,34 +227,45 @@ def make_derivative(expr, dim, fd_order, deriv_order, side, matvec, x0, coeffici
222227
if callable(expand):
223228
expand = expand(dim)
224229

230+
# With a `subdomain`, `expr` is treated as zero outside it: every stencil tap
231+
# is multiplied by the indicator of the point it reads, 1 if `subdomain`
232+
# spans the whole of `dim`. The derivative then extends past `subdomain` by
233+
# the stencil radius
234+
halo_radius = {dim.root: indices.radius} if subdomain is not None else {}
235+
225236
if not expand and indices.expr is not None:
226237
weights = Weights(name='w', dimensions=indices.free_dim,
227238
initvalue=weights, dtype=expr.dtype)
228239

229240
# Inject the StencilDimension
230241
# E.g. `x + i*h_x` into `f(x)` s.t. `f(x + i*h_x)`
231242
expr = expr.shift(dim, indices.expr - dim)
243+
if subdomain is not None:
244+
expr = expr * subdomain.indicator(dim.root, indices.offset(indices.expr))
232245

233246
# Re-evaluate any off-the-grid Functions potentially impacted by the FD
234247
# unless a pure number
235248
with suppress(AttributeError):
236249
expr = expr._evaluate(expand=False)
237250

238251
deriv = DiffDerivative(
239-
expr*weights, {dim: indices.free_dim}, deriv_order=deriv_order
252+
expr*weights, {dim: indices.free_dim}, deriv_order=deriv_order,
253+
halo_radius=halo_radius
240254
)
241255
else:
242256
terms = []
243257
for i, c in zip(indices, weights, strict=True):
244258
# The FD term
245259
term = expr.shift(dim, i - dim) * c
260+
if subdomain is not None:
261+
term = term * subdomain.indicator(dim.root, indices.offset(i))
246262

247263
# Re-evaluate any off-the-grid Functions potentially impacted by the FD
248264
# unless a pure number
249265
with suppress(AttributeError):
250266
term = term.evaluate
251267
terms.append(term)
252268

253-
deriv = EvalDerivative(*terms, base=expr)
269+
deriv = EvalDerivative(*terms, base=expr, halo_radius=halo_radius)
254270

255271
return deriv

0 commit comments

Comments
 (0)