Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 93 additions & 20 deletions devito/finite_differences/derivative.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,7 @@ def _fd_priority(self):

__rargs__ = ('expr', '*dims')
__rkwargs__ = ('side', 'deriv_order', 'fd_order', 'transpose', '_ppsubs',
'x0', 'method', 'weights')
'x0', 'method', 'weights', 'halo', 'subdomain')

def __new__(cls, expr, *dims, **kwargs):
# Validate the input arguments `expr`, `dims` and `deriv_order`
Expand Down Expand Up @@ -161,6 +161,8 @@ def __new__(cls, expr, *dims, **kwargs):
obj._transpose = kwargs.get("transpose", direct)
obj._method = kwargs.get("method", 'FD')
obj._weights = cls._process_weights(**kwargs)
obj._halo = cls._validate_halo(kwargs.get("halo"))
Comment thread
mloubout marked this conversation as resolved.
obj._subdomain = kwargs.get("subdomain")

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

return obj

@staticmethod
def _validate_halo(halo):
"""
Validate `halo`. Only None (read the argument everywhere) and 0 (treat
the argument as zero outside the equation's SubDomain) are supported.
"""
if halo not in (None, 0):
raise ValueError(f"Expected halo=None or halo=0, got halo={halo}")
Comment thread
FabioLuporini marked this conversation as resolved.
return halo

@staticmethod
def _validate_expr(expr):
"""
Expand Down Expand Up @@ -325,6 +337,8 @@ def _process_weights(cls, **kwargs):
def __call__(self, x0=None, fd_order=None, side=None, method=None, **kwargs):
weights = kwargs.get('weights', kwargs.get('w'))
rkw = {}
if 'halo' in kwargs:
rkw['halo'] = kwargs['halo']
if side is not None:
rkw['side'] = side
if method is not None:
Expand Down Expand Up @@ -457,6 +471,26 @@ def side(self):
def transpose(self):
return self._transpose

@property
def halo(self):
"""
None if the argument is read everywhere, 0 if it is treated as zero
outside the SubDomain of the equation the Derivative belongs to.
"""
return self._halo

@property
def subdomain(self):
"""
With halo=0, the SubDomain outside of which the argument is treated as

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

just to be sure, do you actually need both halo and subdomain? or is subdomain just enough (essentially , when NOT None, it encodes the fact the users wants 0-valued taps outside of it) by any chance?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is a good point

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This only gets set at evaluate when it neeeds one. So conceptually it could but it would be very intricated and would need a lot of weird logic to use only one

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I don't understand why that would be the case?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Because the subdomain cannot be constructed without knowing the radius

zero, set upon evaluation within an equation restricted to it.
"""
return self._subdomain

@cached_property
Comment thread
EdCaunt marked this conversation as resolved.
def _has_zero_halo(self):
return self.halo is not None or self.expr._has_zero_halo

@property
def is_TimeDependent(self):
return self.expr.is_TimeDependent
Expand All @@ -481,26 +515,21 @@ def T(self):

return self._rebuild(transpose=adjoint)

def _eval_at(self, func, interp_mode='direct', **kwargs):
def _eval_at(self, func, interp_mode='direct', subdomain=None, **kwargs):
"""
Evaluates the derivative at the location of `func`. It is necessary for staggered
setup where one could have Eq(u(x + h_x/2), v(x).dx)) in which case v(x).dx
has to be computed at x=x + h_x/2.

With halo=0, the argument is treated as zero outside `subdomain`, which
the Derivative records.
"""
# No staggering, don't waste time
if not self.expr.staggered and not func.staggered:
return self
# If an x0 already exists or evaluating at the same function (i.e u = u.dx)
# do not overwrite it
if self.x0 or self.side is not None or func.function is self.expr.function:
return self
# For basic equation of the form f = Derivative(g, ...) we can just
# compare staggering
if self.expr.staggered == func.staggered and self.expr.is_Function:
return self
# Time derivatives are not affected by space staggering
if all(d.is_Time for d in self.dims):
return self
rkw = {}
if subdomain is not None and self.halo is not None:
rkw['subdomain'] = subdomain

if not self._is_relocated(func):
return self._rebuild(**rkw) if rkw else self

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

Expand All @@ -534,7 +563,8 @@ def _eval_at(self, func, interp_mode='direct', **kwargs):
return self._rebuild(self.expr, **rkw)
args = [self.expr.func(*v) for v in mapper.values()]
args.extend([a for a in self.expr.args if a not in self.expr._args_diff])
args = [self._rebuild(a)._eval_at(func, interp_mode=interp_mode, **kwargs)
args = [self._rebuild(a)._eval_at(func, interp_mode=interp_mode,
subdomain=subdomain, **kwargs)
for a in args]
return self.expr.func(*args)
elif self.expr.is_Mul:
Expand All @@ -549,6 +579,25 @@ def _eval_at(self, func, interp_mode='direct', **kwargs):
# the expression as is.
return self._rebuild(self.expr, **rkw)

def _is_relocated(self, func):
"""
True if the Derivative must be evaluated at the location of `func`, e.g.
with `func` and the argument staggered apart.
"""
# No staggering, don't waste time
if not self.expr.staggered and not func.staggered:
return False
# If an x0 already exists or evaluating at the same function (i.e u = u.dx)
# do not overwrite it
if self.x0 or self.side is not None or func.function is self.expr.function:
return False
# For basic equation of the form f = Derivative(g, ...) we can just
# compare staggering
if self.expr.staggered == func.staggered and self.expr.is_Function:
return False
# Time derivatives are not affected by space staggering
return not all(d.is_Time for d in self.dims)

def _evaluate(self, **kwargs):
# Evaluate finite-difference.
# NOTE: `evaluate` and `_eval_fd` split for potential future different
Expand Down Expand Up @@ -581,6 +630,13 @@ def _eval_fd(self, expr, **kwargs):
if expr.is_Add and any(len(indices_at(expr, d)) > 1 for d in self.dims):
return expr.func(*[self._eval_fd(a, **kwargs) for a in expr.args])

# The SubDomain mask read by a halo=0 derivative can't be shifted again
if any(d.halo is not None for d in expr.find(Derivative)):
raise NotImplementedError(
f"{self} differentiates a derivative with halo=0, which is not "
"supported"
)

# Step 1: Evaluate non-derivative x0. We currently enforce a simple 2nd order
# interpolation to avoid very expensive finite differences on top of it
x0_deriv = self._filter_dims(self.x0)
Expand All @@ -598,6 +654,15 @@ def _eval_fd(self, expr, **kwargs):
# otherwise an IndexSum will returned
expand = kwargs.get('expand', True)

# With halo=0, `expr` is treated as zero outside the equation's SubDomain
subdomain = self.subdomain
if subdomain is not None and (subdomain.is_MultiSubDomain or
self.method != 'FD'):
raise NotImplementedError(
f"halo=0 is only supported with method='FD' on a SubDomain, not "
f"with method={self.method} on {subdomain}"
)

# Step 3: Evaluate FD of the new expression
if self.method == 'RSFD':
assert len(self.dims) == 1
Expand All @@ -607,18 +672,26 @@ def _eval_fd(self, expr, **kwargs):
assert self.method == 'FD'
res = cross_derivative(expr, self.dims, self.fd_order, self.deriv_order,
matvec=self.transpose, x0=x0_deriv, expand=expand,
side=self.side, weights=self.weights)
side=self.side, weights=self.weights,
subdomain=subdomain)
else:
assert self.method == 'FD'
res = generic_derivative(expr, self.dims[0], self.fd_order[0],
self.deriv_order[0], weights=self.weights,
side=self.side, matvec=self.transpose,
x0=self.x0, expand=expand)
x0=self.x0, expand=expand,
subdomain=subdomain)

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

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

return res

def _eval_expand_nest(self, **hints):
Expand Down
94 changes: 91 additions & 3 deletions devito/finite_differences/differentiable.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,6 +191,25 @@ def _eval_at(self, func, **kwargs):
for a in self.args # false positive: lambda is invoked in-place
])

@cached_property
def _has_zero_halo(self):
"""True if the expression has derivatives with `halo=0`."""
return any(a._has_zero_halo for a in self._args_diff)

@cached_property
def extent(self):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

IIRC, I think "extent" is for physical quantities ? also, is it possible this should be a private property? I doubt the use needs to be able to query it

"""
Number of points, per root Dimension, by which the evaluated expression
extends past the SubDomain that its `halo=0` derivatives restrict their
argument to: the largest stencil radius of such derivatives. Empty if the
expression has none.

For example, with `S` a SubDomain restricting `x` and an 8th-order `g`,
`g.dx(halo=0)` evaluated in an equation on `S` has an extent of `{x: 4}`:
the equation must iterate over `S` grown by 4 points along `x`.
"""
return merge_extent(a.extent for a in self._args_diff)

def _subs(self, old, new, **hints):
if old == self:
return new
Expand Down Expand Up @@ -329,8 +348,18 @@ def shift(self, dim, shift):
"""
Shift expression by `shift` along the Dimension `dim`.
For example u.shift(x, x.spacing) = u(x + h_x).

Every other space Dimension of the expression sharing the root of `dim`
is shifted too: an expression may mix Functions defined on the Grid,
e.g. `f(x)`, with Functions defined on a SubDomain, e.g. `p(ix)`.
"""
return self._subs(dim, dim + shift)
from devito.symbolics import retrieve_dimensions # noqa

expr = self
for d in retrieve_dimensions(self, mode='unique'):
if d is not dim and dim.root in d._defines and not d.is_NonlinearDerived:
expr = expr._subs(d, d + shift)
return expr._subs(dim, dim + shift)

@property
def laplace(self):
Expand Down Expand Up @@ -530,6 +559,18 @@ def highest_priority(diff_op, candidates=None):
return prio_func


def merge_extent(extents):
"""
Merge the extents `extents`, each a mapping from root Dimension to a number
of points, keeping the largest per Dimension (see `Differentiable.extent`).
"""
extent = {}
for i in extents:
for d, v in i.items():
extent[d] = max(extent.get(d, 0), v)
return frozendict(extent)


class DifferentiableOp(Differentiable):

__sympy_class__ = None
Expand Down Expand Up @@ -630,6 +671,19 @@ def __new__(cls, *args, **kwargs):

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

def _eval_at(self, func, subdomain=None, **kwargs):
"""
Evaluate the sum at the location of `func`.

The derivatives with `halo=0` extend the sum past `subdomain`, so the
other terms are restricted to `subdomain` at the evaluation point.
"""
expr = super()._eval_at(func, subdomain=subdomain, **kwargs)
if subdomain is None or not expr.is_Add or not expr._has_zero_halo:
return expr
halo = {a for a in expr._args_diff if a._has_zero_halo}
return self.func(*[a if a in halo else subdomain.restrict(a) for a in expr.args])


class Mul(DifferentiableOp, sympy.Mul):
__sympy_class__ = sympy.Mul
Expand Down Expand Up @@ -1318,6 +1372,25 @@ def _subs(self, old, new, **hints):

class DiffDerivative(IndexDerivative, DifferentiableOp):

"""
A Derivative evaluated in unexpanded form.

Parameters
----------
extent : dict of {Dimension: int}, optional
For a derivative with `halo=0`, its stencil radius per root Dimension
(see `Differentiable.extent`).
*args, **kwargs
As for IndexDerivative.
"""

__rkwargs__ = IndexDerivative.__rkwargs__ + ('extent',)

def __new__(cls, *args, extent=None, **kwargs):
obj = super().__new__(cls, *args, **kwargs)
obj.extent = frozendict(extent or {})
return obj

def _eval_at(self, func, **kwargs):
# Like EvalDerivative, a DiffDerivative must have already been evaluated
# at a valid x0 and should not be re-evaluated at a different location
Expand All @@ -1331,11 +1404,25 @@ def _eval_at(self, func, **kwargs):

class EvalDerivative(DifferentiableOp, sympy.Add):

"""
A Derivative evaluated in expanded form, as the sum of its stencil taps.

Parameters
----------
*args : expr-like
The stencil taps.
base : expr-like, optional
The expression the derivative was taken of.
extent : dict of {Dimension: int}, optional
For a derivative with `halo=0`, its stencil radius per root Dimension
(see `Differentiable.extent`).
"""

is_commutative = True

__rkwargs__ = ('base',)
__rkwargs__ = ('base', 'extent')

def __new__(cls, *args, base=None, **kwargs):
def __new__(cls, *args, base=None, extent=None, **kwargs):
kwargs['evaluate'] = False

# a+0 -> a
Expand All @@ -1351,6 +1438,7 @@ def __new__(cls, *args, base=None, **kwargs):
# In some rare cases (rebuild?) base may be obj itself
base = base.base
obj.base = base
obj.extent = frozendict(extent or {})
except AttributeError:
# This might happen if e.g. one attempts a (re)construction with
# one sole argument. The (re)constructed EvalDerivative degenerates
Expand Down
Loading
Loading