-
Notifications
You must be signed in to change notification settings - Fork 263
api: Fix mixed subdomain derivative #3035
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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` | ||
|
|
@@ -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")) | ||
| obj._subdomain = kwargs.get("subdomain") | ||
|
|
||
| ppsubs = kwargs.get("subs", kwargs.get("_ppsubs", [])) | ||
| processed = [] | ||
|
|
@@ -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}") | ||
|
FabioLuporini marked this conversation as resolved.
|
||
| return halo | ||
|
|
||
| @staticmethod | ||
| def _validate_expr(expr): | ||
| """ | ||
|
|
@@ -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: | ||
|
|
@@ -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 | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. just to be sure, do you actually need both
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This is a good point
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This only gets set at
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I don't understand why that would be the case?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 | ||
|
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 | ||
|
|
@@ -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 | ||
|
|
@@ -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,) | ||
|
|
||
|
|
@@ -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: | ||
|
|
@@ -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 | ||
|
|
@@ -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) | ||
|
|
@@ -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 | ||
|
|
@@ -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): | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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): | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 | ||
|
|
@@ -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): | ||
|
|
@@ -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 | ||
|
|
@@ -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 | ||
|
|
@@ -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 | ||
|
|
@@ -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 | ||
|
|
@@ -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 | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.