diff --git a/devito/ir/equations/algorithms.py b/devito/ir/equations/algorithms.py index 4d8a02c6f5..de477b0292 100644 --- a/devito/ir/equations/algorithms.py +++ b/devito/ir/equations/algorithms.py @@ -393,7 +393,7 @@ def generate_conditionals(expr, input_expr, ordering): index = d.index if d.condition is not None and expr.has(d): index = index - relational_min(cond, d.parent) - shift = relational_shift(cond, d.parent) + shift = relational_shift(cond, d.parent, d.symbolic_factor) expr = uxreplace(expr, {d: IntDiv(index, d.symbolic_factor) + shift}) return expr, conditionals diff --git a/devito/types/relational.py b/devito/types/relational.py index bf2e625a41..42772e7d3b 100644 --- a/devito/types/relational.py +++ b/devito/types/relational.py @@ -294,30 +294,34 @@ def _(expr, s): return sympy.S.Infinity -def relational_shift(expr, s): +def relational_shift(expr, s, factor=None): """ Infer shift incurred by the expression. Generally only applies when a CondEq is used as it adds a single value. + + When `factor` is provided, the CondEq point only adds a value if it does + not already fall onto the subsampling grid defined by `factor`, as it + would otherwise overlap with an already counted point. """ if not expr.has(s): return 0 - return _relational_shift(expr, s) + return _relational_shift(expr, s, factor) @singledispatch -def _relational_shift(s, expr): +def _relational_shift(expr, s, factor): return 0 @_relational_shift.register(sympy.Or) @_relational_shift.register(sympy.And) -def _(expr, s): - return sum([_relational_shift(e, s) for e in expr.args]) +def _(expr, s, factor): + return sum([_relational_shift(e, s, factor) for e in expr.args]) @_relational_shift.register(sympy.Eq) -def _(expr, s): +def _(expr, s, factor): if isinstance(expr.lhs, sympy.Mod): return 0 @@ -328,4 +332,11 @@ def _(expr, s): except (AttributeError, AssertionError): # Stepping dimension (time), requires shift from devito.symbolics.extended_dtypes import INT - return INT(Ge(*expr.args)) + shift = INT(Ge(*expr.args)) + if factor is None: + return shift + # No shift if the point is already on the subsampling grid + remainder = sympy.Mod(expr.rhs, factor) + if remainder.is_Integer: + return shift if remainder != 0 else 0 + return shift * INT(Ne(remainder, 0)) diff --git a/tests/test_buffering.py b/tests/test_buffering.py index 92c803b7ef..ef2a06c3fb 100644 --- a/tests/test_buffering.py +++ b/tests/test_buffering.py @@ -1128,4 +1128,8 @@ def test_buffering_multi_cond(factor): eq_all.append(Eq(f_all, f, implicit_dims=ctend)) op_all = Operator(eq_all, opt='buffering') op_all.apply(time_m=0, time_M=ntmod-2) - assert np.allclose(f_all.data[:, 11, 11], factor * np.arange(nt)) + expected = factor * np.arange(nt) + if (ntmod - 2) % factor == 0: + # The last sample is on the subsampled grid, hence no extra slot + expected[-2:] = [ntmod - 1, 0] + assert np.allclose(f_all.data[:, 11, 11], expected) diff --git a/tests/test_dimension.py b/tests/test_dimension.py index 27d93088cc..3e28210fe5 100644 --- a/tests/test_dimension.py +++ b/tests/test_dimension.py @@ -7,10 +7,10 @@ from conftest import assert_blocking, assert_structure, opts_tiling, skipif from devito import ( # noqa - Buffer, ConditionalDimension, Constant, CustomDimension, DefaultDimension, Dimension, - Eq, Function, Ge, Grid, Gt, Inc, Le, Lt, Ne, Operator, SpaceDimension, SparseFunction, - SparseTimeFunction, SubDimension, SubDomain, TimeFunction, configuration, dimensions, - floor, norm, sin, sum, switchconfig + Buffer, CondEq, ConditionalDimension, Constant, CustomDimension, DefaultDimension, + Dimension, Eq, Function, Ge, Grid, Gt, Inc, Le, Lt, Ne, Operator, SpaceDimension, + SparseFunction, SparseTimeFunction, SubDimension, SubDomain, TimeFunction, + configuration, dimensions, floor, norm, sin, sum, switchconfig ) from devito.exceptions import InvalidArgument from devito.ir import SymbolRegistry @@ -2074,6 +2074,39 @@ def test_factor_and_condition(self): for t in range(buffer_size): assert np.all(usaved.data[t] == t*factor + bounds[0] - 1) + @pytest.mark.parametrize('factor', [1, 2, 3]) + @pytest.mark.parametrize('relation', [Or, 'strict']) + def test_factor_and_condeq_overlap(self, factor, relation): + """ + A snapshot at `t == nt - 1` via an implicit CondEq ConditionalDimension + only adds an extra slot when `nt - 1` is not on the subsampling grid. + Otherwise it must write the slot of the subsampled point, rather than + shift past it (and out of bounds for `factor == 1`). + """ + grid = Grid(shape=(4, 4)) + time = grid.time_dim + + nt = 10 + last = nt - 1 + nsave = (last + factor - 1) // factor + 1 + + ctsnap = ConditionalDimension(name='ctsnap', parent=time, + condition=CondEq(time, last), + relation=relation) + ct0 = ConditionalDimension(name='ct0', parent=time, factor=factor, + relation=Or) + + u = TimeFunction(name='u', grid=grid, save=nt) + usave = TimeFunction(name='usave', grid=grid, time_dim=ct0, save=nsave) + u.data[:] = np.arange(nt).reshape(nt, 1, 1) + + Operator(Eq(usave, u))(time_m=0, time_M=last - 1) + Operator(Eq(usave, u, implicit_dims=ctsnap))(time_m=last, time_M=last) + + expected = factor * np.arange(nsave) + expected[-1] = last + assert np.all(usave.data[:, 1, 1] == expected) + def test_blocking_w_guard(self): grid = Grid(shape=(8, 8, 8)) x, y, z = grid.dimensions