-
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 expr._has_zero_halo: | ||
| 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,29 @@ 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`, the stencil taps above are zeroed outside the SubDomain | ||
| # along `self.dims` only. Along any other Dimension, the argument is read | ||
| # at the evaluation point, which other `halo=0` derivatives may push | ||
| # outside the SubDomain. E.g., in `f.dx(halo=0) + f.dy(halo=0)`, the | ||
| # equation iterates over the SubDomain grown along both `x` and `y`, so | ||
| # `f.dx` must also be zeroed where `y` lies outside the SubDomain | ||
| 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): | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.