-
Notifications
You must be signed in to change notification settings - Fork 257
api: fix handling of multiple conditions for buffering #2850
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
d696748
e6fd96b
d6565a5
0aeef75
60f7897
0771802
05aa7af
09ae281
f92847e
4176610
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 |
|---|---|---|
|
|
@@ -13,7 +13,8 @@ | |
| from sympy.logic.boolalg import BooleanFunction | ||
|
|
||
| from devito.ir.support.space import Forward, IterationDirection | ||
| from devito.symbolics import CondEq, CondNe, search | ||
| from devito.symbolics import CondEq, CondNe, IntDiv, search | ||
| from devito.symbolics.manipulation import _uxreplace_handle, _uxreplace_registry | ||
| from devito.tools import Pickable, as_tuple, frozendict, split | ||
| from devito.types import Dimension, LocalObject | ||
|
|
||
|
|
@@ -64,7 +65,14 @@ class GuardFactor(Guard, CondEq, Pickable): | |
|
|
||
| __rargs__ = ('d',) | ||
|
|
||
| def __new__(cls, d, **kwargs): | ||
| def __new__(cls, *args, **kwargs): | ||
| if len(args) != 1: | ||
| # Reconstruction with relational args (e.g. via sympy `_subs`): the | ||
| # factor semantics no longer hold, so degrade to a plain relational | ||
| base = CondNe if issubclass(cls, CondNe) else CondEq | ||
| return base(*args, **kwargs) | ||
|
|
||
| d, = args | ||
| assert d.is_Conditional | ||
|
|
||
| obj = super().__new__(cls, d.parent % d.symbolic_factor, 0) | ||
|
|
@@ -138,54 +146,54 @@ class BaseGuardBoundNext(Guard, Pickable): | |
| given `direction`. | ||
| """ | ||
|
|
||
| __rargs__ = ('d', 'direction') | ||
| __rargs__ = ('d', 'index', 'direction') | ||
| __rkwargs__ = ('d_min', 'd_max') | ||
|
|
||
| def __new__(cls, d, direction, **kwargs): | ||
| def __new__(cls, d, index, direction, | ||
| d_min=None, d_max=None, **kwargs): | ||
| assert isinstance(d, Dimension) | ||
| assert isinstance(direction, IterationDirection) | ||
|
|
||
| if direction == Forward: | ||
| p0 = d.root | ||
| p1 = d.root.symbolic_max | ||
| # Always take the next index in the iteration direction | ||
| next_index = eval_next_index(index, d, direction) | ||
|
|
||
| if d.is_Conditional: | ||
| v = d.symbolic_factor | ||
| # Round `p0 + 1` up to the nearest multiple of `v` | ||
| p0 = Mul((((p0 + 1) + v - 1) / v), v, evaluate=False) | ||
| else: | ||
| p0 = p0 + 1 | ||
| # The direction might be forward but accessing c - d | ||
| # making the access backward w.r.t | ||
| # Update direction according to access direction for valid guard | ||
| if index.has(-d): | ||
| direction = -direction | ||
|
|
||
| if direction == Forward: | ||
| p0 = next_index | ||
| p1 = d_max or d.root.symbolic_max | ||
| else: | ||
| p0 = d.root.symbolic_min | ||
| p1 = d.root | ||
|
|
||
| if d.is_Conditional: | ||
| v = d.symbolic_factor | ||
| # Round `p1 - 1` down to the nearest sub-multiple of `v` | ||
| # NOTE: we use ABS to make sure we handle negative values properly. | ||
| # Once `p1 - 1` is negative (e.g. `iteration=time - 1` and `time=0`), | ||
| # as long as we get a negative number, rather than 0 and even if it's | ||
| # not `-v`, we're good | ||
| p1 = (p1 - 1) - abs(p1 - 1) % v | ||
| else: | ||
| p1 = p1 - 1 | ||
| p0 = d_min if d_min is not None else d.root.symbolic_min | ||
| p1 = next_index | ||
|
|
||
| try: | ||
| if cls.__base__._eval_relation(p0, p1) is true: | ||
| return None | ||
| except TypeError: | ||
| pass | ||
|
|
||
| return cls._new(p0, p1, d, index, direction, d_min=d_min, d_max=d_max, **kwargs) | ||
|
|
||
| @classmethod | ||
| def _new(cls, p0, p1, d, index, direction, d_min=None, d_max=None): | ||
|
|
||
| obj = super().__new__(cls, p0, p1, evaluate=False) | ||
|
|
||
| obj.d = d | ||
| obj.direction = direction | ||
| obj.index = index | ||
| obj.d_min = d_min | ||
|
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. these should probably be
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. also for homogeneity just
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. Bit of a pain to make them args, there is places that don't use them |
||
| obj.d_max = d_max | ||
|
|
||
| return obj | ||
|
|
||
| @property | ||
| def _args_rebuild(self): | ||
| return (self.d, self.direction) | ||
| return (self.d, self.index, self.direction) | ||
|
|
||
|
|
||
| class GuardBoundNextLe(BaseGuardBoundNext, Le): | ||
|
|
@@ -272,6 +280,15 @@ class Guards(frozendict): | |
| def get(self, d, v=true): | ||
| return super().get(d, v) | ||
|
|
||
| def has(self, d, cls): | ||
| """ | ||
| True if the guard registered for `d` contains an instance of `cls`. | ||
| """ | ||
| g = super().get(d) | ||
| if g is None: | ||
| return False | ||
| return g.has(cls) | ||
|
|
||
| def _reuse_if_untouched(self, mapper): | ||
| return self if mapper == self else Guards(mapper) | ||
|
|
||
|
|
@@ -569,3 +586,49 @@ def pairwise_or(*guards): | |
| pass | ||
|
|
||
| return guard | ||
|
|
||
|
|
||
| _uxreplace_registry.register(BaseGuardBoundNext) | ||
|
|
||
|
|
||
| @_uxreplace_handle.register(BaseGuardBoundNext) | ||
| def _(expr, args, kwargs): | ||
| p0, p1 = args | ||
| return expr._new(p0, p1, expr.d, expr.index, expr.direction, | ||
| **kwargs) | ||
|
|
||
|
|
||
| @singledispatch | ||
| def eval_next_index(expr, dim, dir): | ||
| """ | ||
| Evaluate `expr` at the next iteration point along `dim` in the given | ||
| `dir`-ection. The "next" point is obtained by substituting `dim` with | ||
| `dim + 1` for `Forward` and `dim - 1` for `Backward`. | ||
|
|
||
| For `IntDiv` expressions encoding subsampling (`dim.root // factor`), | ||
| the result is rounded to the next valid coarse-grained slot. | ||
| """ | ||
| if dir == Forward: | ||
| return expr._subs(dim, dim + 1) | ||
| else: | ||
| return expr._subs(dim, dim - 1) | ||
|
|
||
|
|
||
| @eval_next_index.register(Expr) | ||
| def _(expr, dim, dir): | ||
| if not expr.args: | ||
| if dir == Forward: | ||
| return expr._subs(dim, dim + 1) | ||
| else: | ||
| return expr._subs(dim, dim - 1) | ||
| return expr.func(*[eval_next_index(a, dim, dir) for a in expr.args]) | ||
|
|
||
|
|
||
| @eval_next_index.register(IntDiv) | ||
| def _(expr, dim, dir): | ||
| v = dim.symbolic_factor | ||
| p0 = dim.root | ||
| if dir == Forward: | ||
| return Mul((((p0 + 1) + v - 1) / v), v, evaluate=False) | ||
| else: | ||
| return (p0 - 1) - abs(p0 - 1) % v | ||
Uh oh!
There was an error while loading. Please reload this page.