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
2 changes: 1 addition & 1 deletion devito/ir/equations/algorithms.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
25 changes: 18 additions & 7 deletions devito/types/relational.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

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.

Adding an example to this docstring would be useful. I figured it out, but it wasn't super obvious

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

Expand All @@ -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))
6 changes: 5 additions & 1 deletion tests/test_buffering.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
41 changes: 37 additions & 4 deletions tests/test_dimension.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Loading