diff --git a/devito/finite_differences/differentiable.py b/devito/finite_differences/differentiable.py index 57f0db8e08..d86734a454 100644 --- a/devito/finite_differences/differentiable.py +++ b/devito/finite_differences/differentiable.py @@ -21,8 +21,8 @@ from devito.finite_differences.tools import coeff_priority, make_shift_x0 from devito.logger import warning from devito.tools import ( - Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype, - is_integer, is_number, memoized_func, split + Pickable, Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict, + infer_dtype, is_integer, is_number, memoized_func, split ) from devito.types import Array, DimensionTuple, Evaluable, StencilDimension from devito.types.basic import AbstractFunction, Indexed @@ -35,6 +35,7 @@ 'Imag', 'IndexDerivative', 'IndexDerivativeProperty', + 'LocalSum', 'Real', 'Weights', ] @@ -940,13 +941,106 @@ def _evaluate(self, **kwargs): terms.append(expr.xreplace(mapper)) return sum(terms) + @property + def bound_symbols(self): + return set(self.dimensions) + @property def free_symbols(self): - return super().free_symbols - set(self.dimensions) + return super().free_symbols - self.bound_symbols func = DifferentiableOp._rebuild +class LocalSum(IndexSum, Pickable): + + """ + A zero-initialized sum over guarded local dimensions. + + `cdims` are guarded ConditionalDimensions, retained with their original + parents and conditions. `dimensions` exposes the parent iteration dimensions. + Masked points contribute zero. The sum remains symbolic until Cluster lowering + chooses its implementation. + + Examples + -------- + For bilinear interpolation, `posx` and `posy` are the grid indices of sparse + point `p`, and `wx` and `wy` hold its interpolation weights:: + + i = CustomDimension('i', 0, 1, 2) + j = CustomDimension('j', 0, 1, 2) + ci = ConditionalDimension('i', i, indirect=True, + condition=And(posx + i >= x_m, posx + i <= x_M)) + cj = ConditionalDimension('j', j, indirect=True, + condition=And(posy + j >= y_m, posy + j <= y_M)) + value = LocalSum( + wx[p, ci]*wy[p, cj]*f[posx + ci, posy + cj], + cdims=(ci, cj) + ) + Eq(rcv[p], value) + + The scalar lowering has the following semantics (pseudocode):: + + acc = 0 + for i in range(2): + for j in range(2): + if x_m <= posx + i <= x_M and y_m <= posy + j <= y_M: + acc += wx[p, i]*wy[p, j]*f[posx + i, posy + j] + rcv[p] = acc + + The guarded indices and their parents are local to the sum; `p` remains an + outer iteration dimension. + If every tap is masked, `rcv[p]` receives zero. + """ + + __rargs__ = ('expr',) + __rkwargs__ = ('cdims', 'dtype') + + def __new__(cls, expr, cdims=(), dtype=None, **kwargs): + obj = sympy.Expr.__new__(cls, expr) + + obj._expr = expr + obj._cdims = as_tuple(cdims) + obj._dtype = dtype + + return obj + + def _hashable_content(self): + return super()._hashable_content() + (self.cdims, self.dtype) + + @property + def cdims(self): + return self._cdims + + @cached_property + def dtype(self): + if self._dtype is None: + return extract_dtype(self.expr) + return self._dtype + + @cached_property + def dimensions(self): + return tuple(d.parent for d in self.cdims) + + @cached_property + def conditionals(self): + return frozendict({d: d.condition for d in self.cdims}) + + @property + def bound_symbols(self): + return super().bound_symbols | set(self.cdims) + + @property + def free_symbols(self): + symbols = self.expr.free_symbols.union(*[d.free_symbols for d in self.cdims]) + return symbols - self.bound_symbols + + def _evaluate(self, **kwargs): + return self._rebuild(*self._evaluate_args(**kwargs)) + + __reduce_ex__ = Pickable.__reduce_ex__ + + class WeightsIndexed(Indexed): @property diff --git a/devito/ir/clusters/cluster.py b/devito/ir/clusters/cluster.py index 2919e63642..bba54fbafc 100644 --- a/devito/ir/clusters/cluster.py +++ b/devito/ir/clusters/cluster.py @@ -148,6 +148,13 @@ def dist_dimensions(self): ret.update(f._dist_dimensions) return frozenset(ret) + @cached_property + def local_sums(self): + """ + The local sums in equation order, retaining occurrences across equations. + """ + return tuple(s for e in self.exprs for s in e.local_sums) + @cached_property def scope(self): return Scope(self.exprs) @@ -674,6 +681,10 @@ def rebuild(self, **kwargs): def exprs(self): return flatten(c.exprs for c in self) + @cached_property + def local_sums(self): + return tuple(s for c in self for s in c.local_sums) + @cached_property def scope(self): return Scope(exprs=self.exprs) diff --git a/devito/ir/equations/algorithms.py b/devito/ir/equations/algorithms.py index f5eb0e6535..4d8a02c6f5 100644 --- a/devito/ir/equations/algorithms.py +++ b/devito/ir/equations/algorithms.py @@ -3,7 +3,7 @@ from devito.data.allocators import DataReference from devito.finite_differences.differentiable import diff2sympy -from devito.ir.support import GuardFactor +from devito.ir.support import GuardFactor, bounded from devito.logger import warning from devito.symbolics import ( IntDiv, retrieve_dimensions, retrieve_functions, retrieve_indexed, uxreplace @@ -28,12 +28,17 @@ def dimension_sort(expr): appear within Indexeds. """ + # Bound Dimensions do not impact the order of the enclosing iteration space + bound = bounded(expr) + def handle_indexed(indexed): relation = [] for i in indexed.indices: try: # Assume it's an AffineIndexAccessFunction... - relation.append(i.d) + # It may contain only a scalar offset and bound stencil indices + if i.d: + relation.append(i.d) except AttributeError: # It's not! Maybe there are some nested Indexeds (e.g., the # situation is A[B[i]]) @@ -46,9 +51,9 @@ def handle_indexed(indexed): # what the user is attempting to do relation.extend(filter_sorted(i.atoms(Dimension))) - # StencilDimensions are lowered subsequently through special compiler + # Bound Dimensions are lowered subsequently through special compiler # passes, so they can be ignored here - relation = tuple(d for d in relation if not d.is_Stencil) + relation = tuple(d for d in relation if d not in bound) return relation @@ -61,7 +66,7 @@ def handle_indexed(indexed): relations.add(expr.implicit_dims) # Add in leftover free dimensions (not an Indexed' index) - extra = set(retrieve_dimensions(expr, deep=True)) + extra = set(retrieve_dimensions(expr, deep=True)) - bound # Add in pure data dimensions (e.g., those accessed only via explicit values, # such as A[3]) diff --git a/devito/ir/equations/equation.py b/devito/ir/equations/equation.py index cb7588f72c..3139cb3354 100644 --- a/devito/ir/equations/equation.py +++ b/devito/ir/equations/equation.py @@ -4,14 +4,15 @@ import numpy as np import sympy -from devito.finite_differences.differentiable import diff2sympy +from devito.finite_differences.differentiable import LocalSum, diff2sympy from devito.ir.equations.algorithms import dimension_sort, generate_conditionals from devito.ir.support import ( - Interval, IntervalGroup, IterationSpace, Stencil, detect_accesses + Interval, IntervalGroup, IterationSpace, Stencil, bounded, detect_accesses ) -from devito.symbolics import limits_mapper, retrieve_accesses +from devito.symbolics import limits_mapper, retrieve_accesses, search from devito.tools import ( - Pickable, Tag, as_hashable, filter_sorted, frozendict, reuse_if_unchanged + Pickable, Tag, as_hashable, filter_ordered, filter_sorted, frozendict, + reuse_if_unchanged ) from devito.types import Eq, Inc, ReduceMax, ReduceMin, ReduceMinMax @@ -50,6 +51,11 @@ def ispace(self): def dimensions(self): return set(self.ispace.dimensions) + @cached_property + def local_sums(self): + """The local sums in dependency order, with nested sums first.""" + return tuple(filter_ordered(search(self, LocalSum, mode='all'))) + @property def implicit_dims(self): return self._implicit_dims @@ -273,11 +279,15 @@ def __new__(cls, *args, **kwargs): # Analyze the expression accesses = detect_accesses(expr) dimensions = Stencil.union(*accesses.values()) + bound = bounded(expr) # Separate out the SubIterators from the main iteration Dimensions, that # is those which define an actual iteration space iterators = {} for d in dimensions: + if d in bound: + # Local sum dimensions belong to the sum's own iteration space + continue if d.is_SubIterator: iterators.setdefault(d.root, set()).add(d) elif d.is_Conditional: diff --git a/devito/ir/support/utils.py b/devito/ir/support/utils.py index 5f2ee39af7..67844287a9 100644 --- a/devito/ir/support/utils.py +++ b/devito/ir/support/utils.py @@ -2,7 +2,7 @@ from contextlib import suppress from itertools import product -from devito.finite_differences import IndexDerivative +from devito.finite_differences.differentiable import IndexDerivative, IndexSum from devito.symbolics import retrieve_indexed, search from devito.tools import DefaultOrderedDict, as_tuple, filter_sorted, split from devito.types import ( @@ -13,6 +13,7 @@ 'AccessMode', 'IMask', 'Stencil', + 'bounded', 'detect_accesses', 'erange', 'extrema', @@ -231,7 +232,18 @@ def pull_dims(exprs, flag=True): return dims -# *** Utility functions for expressions that potentially contain StencilDimensions +# *** Utility functions for bound and unbound Dimensions + + +def bounded(expr): + """ + Retrieve all Dimensions bound by symbolic sums in `expr`. + """ + sums = search(expr, IndexSum, mode='unique', deep=True) + dims = set().union(*(i.bound_symbols for i in sums)) + + return dims - expr.free_symbols + def unbounded(expr): """ diff --git a/devito/operations/interpolators.py b/devito/operations/interpolators.py index a7fa4d76be..ff4e14b1de 100644 --- a/devito/operations/interpolators.py +++ b/devito/operations/interpolators.py @@ -11,14 +11,14 @@ except ImportError: from numpy import i0 -from devito.finite_differences.differentiable import Mul +from devito.finite_differences.differentiable import LocalSum, Mul from devito.finite_differences.elementary import floor from devito.logger import warning from devito.symbolics import INT, retrieve_function_carriers, retrieve_functions from devito.tools import ( Pickable, as_fp64_decimal, as_list, as_tuple, filter_ordered, flatten, memoized_meth ) -from devito.types import CustomDimension, Eq, Evaluable, Inc, SubFunction, Symbol +from devito.types import CustomDimension, Eq, Evaluable, Inc, SubFunction from devito.types.utils import DimensionTuple __all__ = ['LinearInterpolator', 'NearestInterpolator', @@ -453,19 +453,14 @@ def _interp_idx(self, variables, implicit_dims=None, subdomain=None, return idx_subs, temps - def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None): + def _local_accumulator(self, expr, idx_subs, subdomain=None): """ - Generate a local accumulator for the interpolation/injection operation. + Represent the local sum of weighted interpolation contributions. """ - # Accumulate point-wise contributions into a temporary - rhs = Symbol(name=f'sum{self.sfunction.name}', dtype=self.sfunction.dtype) - summands = [Eq(rhs, 0., implicit_dims=implicit_dims)] - # Substitute coordinate base symbols into the interpolation coefficients weights = self._weights(subdomain=subdomain) - summands.extend([Inc(rhs, (weights * expr).xreplace(idx_subs), - implicit_dims=implicit_dims)]) - - return summands, rhs + rdims = self._rdim(subdomain=subdomain) + summand = (weights * expr).xreplace(idx_subs) + return LocalSum(summand, cdims=rdims, dtype=self.sfunction.dtype) @check_radius @check_coords @@ -539,16 +534,13 @@ def _interpolate(self, expr, increment=False, self_subs=None, implicit_dims=None idx_subs, temps = self._interp_idx(variables, implicit_dims=implicit_dims, subdomain=subdomain) - # Local scalar for accumulation over radius - summands, rhs = self._local_accumulator(expr, idx_subs, - implicit_dims=implicit_dims, - subdomain=subdomain) + rhs = self._local_accumulator(expr, idx_subs, subdomain=subdomain) # Write/Incr `self` lhs = self.sfunction.subs(self_subs) ecls = Inc if increment else Eq last = [ecls(lhs, rhs, implicit_dims=implicit_dims)] - return temps + summands + last + return temps + last def _inject(self, field, expr, increment=True, implicit_dims=None): """ @@ -785,8 +777,8 @@ class NearestInterpolator(LinearInterpolator): _name = 'nearest' - def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None): - return [], expr.xreplace(idx_subs) + def _local_accumulator(self, expr, idx_subs, subdomain=None): + return expr.xreplace(idx_subs) @memoized_meth def _rdim(self, subdomain=None, shifts=None): diff --git a/devito/operator/operator.py b/devito/operator/operator.py index 174714d0bf..c97870f075 100644 --- a/devito/operator/operator.py +++ b/devito/operator/operator.py @@ -33,7 +33,8 @@ from devito.parameters import configuration from devito.passes import ( Graph, error_mapper, finalize_args, generate_implicit, generate_macros, is_on_device, - lower_dtypes, lower_index_derivatives, minimize_symbols, optimize_pows, unevaluate + lower_dtypes, lower_index_derivatives, lower_local_sums, minimize_symbols, + optimize_pows, unevaluate ) from devito.symbolics import estimate_cost, subs_op_args from devito.tools import ( @@ -423,6 +424,7 @@ def _lower_clusters(cls, expressions, profiler=None, **kwargs): clusters = generate_implicit(clusters) # Lower all remaining high order symbolic objects + clusters = lower_local_sums(clusters, **kwargs) clusters = lower_index_derivatives(clusters, **kwargs) # Turn pows into multiplications. This must happen as late as possible diff --git a/devito/passes/clusters/__init__.py b/devito/passes/clusters/__init__.py index e27d2b755d..bcdb8a4bd6 100644 --- a/devito/passes/clusters/__init__.py +++ b/devito/passes/clusters/__init__.py @@ -9,4 +9,5 @@ from .implicit import * # noqa from .misc import * # noqa from .derivatives import * # noqa +from .localsum import * # noqa from .unevaluate import * # noqa diff --git a/devito/passes/clusters/cse.py b/devito/passes/clusters/cse.py index 1cb2159890..87770d4039 100644 --- a/devito/passes/clusters/cse.py +++ b/devito/passes/clusters/cse.py @@ -11,7 +11,7 @@ # Moved in 1.13 from sympy.core.basic import ordering_of_classes -from devito.finite_differences.differentiable import IndexDerivative +from devito.finite_differences.differentiable import IndexSum from devito.ir import Cluster, Scope, cluster_pass from devito.symbolics import ( DefFunction, Reserved, estimate_cost, q_leaf, q_terminal, search @@ -427,7 +427,7 @@ def _(expr): return {} -@_catch.register(IndexDerivative) +@_catch.register(IndexSum) def _(expr): """ Handler for symbol-binding objects. There can be many of them and therefore diff --git a/devito/passes/clusters/localsum.py b/devito/passes/clusters/localsum.py new file mode 100644 index 0000000000..576f05948e --- /dev/null +++ b/devito/passes/clusters/localsum.py @@ -0,0 +1,75 @@ +from devito.ir import ClusterizedEq, Interval, IterationSpace +from devito.symbolics import uxreplace +from devito.tools import timed_pass +from devito.types import Eq, Inc, Temp + +__all__ = ['lower_local_sum', 'lower_local_sums'] + + +@timed_pass() +def lower_local_sums(clusters, sregistry=None, **kwargs): + """ + Lower LocalSums into private initializers and guarded accumulations. + + For example, using the interpolation's radius dimension and guard:: + + i = CustomDimension('i', 0, 1, 2) + ci = ConditionalDimension('i', i, indirect=True, + condition=And(pos + i >= x_m, pos + i <= x_M)) + Eq(rcv[p], LocalSum(w[p, ci]*f[pos + ci], cdims=(ci,))) + + becomes the following computation at each sparse point `p` (pseudocode):: + + sum0 = 0 + for i in range(2): + if x_m <= pos + i <= x_M: + sum0 += w[p, i]*f[pos + i] + rcv[p] = sum0 + + The initializer and result stay outside the tap guard, so even a fully + masked stencil assigns zero to `rcv[p]`. + """ + processed = [] + for c in clusters: + if not c.local_sums: + processed.append(c) + continue + + for e in c.exprs: + subs = {} + for reduction in e.local_sums: + init, update, value = lower_local_sum( + c, uxreplace(reduction, subs), sregistry + ) + processed.extend([init, update]) + subs[reduction] = value + + expr = uxreplace(e, subs) + processed.append(c.rebuild(exprs=[expr])) + + return processed + + +def lower_local_sum(cluster, reduction, sregistry): + """ + Construct the private initializer and guarded accumulation for one sum. + """ + value = Temp(name=sregistry.make_name(prefix='sum'), dtype=reduction.dtype) + + dims = reduction.dimensions + inner = IterationSpace([Interval(d) for d in dims]) + ispace = IterationSpace.union( + cluster.ispace, inner, relations=(cluster.ispace.itdims + dims,) + ) + + init = cluster.rebuild(exprs=[Eq(value, 0)]) + # Attach the interpolation's original guards to the accumulation only + expr = ClusterizedEq( + Inc(value, reduction.expr), ispace=ispace, conditionals=reduction.conditionals + ) + + properties = cluster.properties.sequentialize(dims) + + update = cluster.rebuild(exprs=expr, ispace=ispace, properties=properties) + + return init, update, value diff --git a/devito/symbolics/inspection.py b/devito/symbolics/inspection.py index e19cd89ad1..57d536bf1c 100644 --- a/devito/symbolics/inspection.py +++ b/devito/symbolics/inspection.py @@ -7,7 +7,7 @@ from sympy.core.numbers import ImaginaryUnit from devito.finite_differences import Derivative -from devito.finite_differences.differentiable import IndexDerivative +from devito.finite_differences.differentiable import IndexDerivative, IndexSum from devito.logger import warning from devito.symbolics.extended_dtypes import INT from devito.symbolics.extended_sympy import CallFromPointer, Cast, DefFunction, Reserved @@ -265,7 +265,7 @@ def _(expr, estimate, seen): return _estimate_cost(expr._evaluate(expand=False), estimate, seen) -@_estimate_cost.register(IndexDerivative) +@_estimate_cost.register(IndexSum) @dont_count_if_seen def _(expr, estimate, seen): flops, _ = _estimate_cost(expr.expr, estimate, seen) @@ -275,7 +275,7 @@ def _(expr, estimate, seen): # To be multiplied by the number of points this index sum implicitly # iterates over - flops *= prod(i._size for i in expr.dimensions) + flops *= prod(i.symbolic_size for i in expr.dimensions) return flops, False diff --git a/devito/symbolics/manipulation.py b/devito/symbolics/manipulation.py index 8de6389825..b545290a15 100644 --- a/devito/symbolics/manipulation.py +++ b/devito/symbolics/manipulation.py @@ -4,10 +4,13 @@ import numpy as np from sympy import Add, Max, Min, Mul, Pow, S, SympifyError, Tuple, sympify +from sympy.core import Basic as SympyBasic from sympy.core.add import _addsort from sympy.core.mul import _mulsort -from devito.finite_differences.differentiable import EvalDerivative, IndexDerivative +from devito.finite_differences.differentiable import ( + EvalDerivative, IndexDerivative, LocalSum +) from devito.symbolics.extended_dtypes import LONG from devito.symbolics.extended_sympy import DefFunction, rfunc from devito.symbolics.queries import q_leaf @@ -72,6 +75,7 @@ def uxreplace(expr, rule): return _uxreplace(expr, rule)[0] +@singledispatch def _uxreplace(expr, rule): if expr in rule: v = rule[expr] @@ -113,12 +117,27 @@ def _uxreplace(expr, rule): return expr, False +@_uxreplace.register(LocalSum) +def _(expr, rule): + # Only the sum owns these guards; ordinary ConditionalDimensions stay atomic + mapper = {} + for d in expr.cdims: + kwargs = {i: getattr(d, i) for i in d.__rkwargs__} + kwargs, changed = _uxreplace_dispatch(kwargs, rule) + if changed: + mapper[d] = d._rebuild(**kwargs) + + # Apply the same dimension substitutions to the summand and its bindings + return _uxreplace.dispatch(object)(expr, {**mapper, **rule}) + + @singledispatch def _uxreplace_dispatch(unknown, rule): return unknown, False @_uxreplace_dispatch.register(Basic) +@_uxreplace_dispatch.register(SympyBasic) def _(expr, rule): return _uxreplace(expr, rule) @@ -174,6 +193,7 @@ def _(expr, args, kwargs): @_uxreplace_handle.register(Add) def _(expr, args, kwargs): + args = [i for i in args if i != 0] if all(i.is_commutative for i in args): _addsort(args) _eval_numbers(expr, args) @@ -232,6 +252,7 @@ def dispatchable(self, obj): _uxreplace_registry.register(Eq) _uxreplace_registry.register(DefFunction) _uxreplace_registry.register(ComponentAccess) +_uxreplace_registry.register(LocalSum) class Uxmapper(dict): diff --git a/examples/userapi/06_sparse_operations.ipynb b/examples/userapi/06_sparse_operations.ipynb index 286ad0297b..66aabe1d35 100644 --- a/examples/userapi/06_sparse_operations.ipynb +++ b/examples/userapi/06_sparse_operations.ipynb @@ -279,9 +279,7 @@ "text": [ "Eq(posx, s_gp(p_s, 0))\n", "Eq(posy, s_gp(p_s, 1))\n", - "Eq(sums, 0.0)\n", - "Inc(sums, s_wx(p_s, rp_sx)*s_wy(p_s, rp_sy)*f(t, rp_sx + posx, rp_sy + posy))\n", - "Eq(s(time, p_s), sums)\n" + "Eq(s(time, p_s), LocalSum(s_wx(p_s, rp_sx)*s_wy(p_s, rp_sy)*f(t, rp_sx + posx, rp_sy + posy), (rp_sx, rp_sy)))\n" ] } ], @@ -484,9 +482,7 @@ "text": [ "Eq(posx, s_gp(p_s, 0))\n", "Eq(posy, s_gp(p_s, 1))\n", - "Eq(sums, 0.0)\n", - "Inc(sums, wsincrp_sx(p_s, rp_sx + 3)*wsincrp_sy(p_s, rp_sy + 3)*f(t, rp_sx + posx, rp_sy + posy))\n", - "Eq(s(time, p_s), sums)\n" + "Eq(s(time, p_s), LocalSum(wsincrp_sx(p_s, rp_sx + 3)*wsincrp_sy(p_s, rp_sy + 3)*f(t, rp_sx + posx, rp_sy + posy), (rp_sx, rp_sy)))\n" ] } ], diff --git a/tests/test_interpolation.py b/tests/test_interpolation.py index af1f1954a3..f2ba3430cd 100644 --- a/tests/test_interpolation.py +++ b/tests/test_interpolation.py @@ -2,7 +2,7 @@ import pytest import scipy.sparse from numpy import floor, sin -from sympy import Float +from sympy import Float, Integer from conftest import assert_structure from devito import ( @@ -11,10 +11,14 @@ SparseFunction, SparseTimeFunction, SubDomain, TimeFunction, VectorFunction, switchconfig ) +from devito.finite_differences import LocalSum +from devito.ir import FindSymbols, LoweredEq from devito.operations.interpolators import ( LinearInterpolator, SincInterpolator, _cell_indices ) +from devito.symbolics import uxreplace from devito.tools import as_tuple +from devito.types import Temp from examples.seismic import ( AcquisitionGeometry, Receiver, RickerSource, TimeAxis, demo_model ) @@ -26,6 +30,77 @@ class SparseFirst(SparseFunction): _sparse_position = 0 +class TestLocalSum: + + @pytest.mark.parametrize('interpolation,r', [('linear', 1), ('sinc', 4)]) + def test_symbolic(self, interpolation, r): + grid = Grid(shape=(17, 19)) + f = Function(name='f', grid=grid, space_order=8) + rcv = SparseFunction(name='rcv', grid=grid, npoint=3, + interpolation=interpolation, r=r) + expr = rcv.interpolate(f).evaluate[-1] + reduction = expr.rhs + + assert isinstance(reduction, LocalSum) + assert tuple(d.symbolic_size for d in reduction.dimensions) == (2*r, 2*r) + assert reduction.free_symbols.isdisjoint(reduction.bound_symbols) + lowered = LoweredEq(expr) + assert lowered.ispace.itdims == (rcv._sparse_dim,) + assert not lowered.conditionals + rdims = rcv.interpolator._rdim() + assert reduction.cdims == rdims + assert reduction.conditionals == {d: d.condition for d in rdims} + assert reduction.evaluate == reduction + assert reduction.func(reduction.expr) == reduction + + # Bounds held only in the guards must participate in substitutions + x, _ = grid.dimensions + mapper = {x.symbolic_max: Integer(10)} + replaced = uxreplace(reduction, mapper) + assert str(replaced.expr) == str(reduction.expr) + assert replaced.conditionals != reduction.conditionals + assert LoweredEq(expr._rebuild(rhs=replaced)).ispace == lowered.ispace + assert x.symbolic_max not in replaced.free_symbols + + @pytest.mark.parametrize('increment', [False, True]) + @pytest.mark.parametrize('opt', ['noop', 'advanced']) + def test_masked(self, increment, opt): + grid = Grid(shape=(17,), extent=(16.,)) + f = Function(name='f', grid=grid) + rcv = SparseFunction(name='rcv', grid=grid, npoint=3) + f.data_with_halo[:] = 2.5 + rcv.coordinates.data[:, 0] = [0., 8.5, 16.] + rcv.data[:] = 7. + op = Operator(rcv.interpolate(f, increment=increment), + name='MaskedSparseSum', opt=opt) + op.apply(x_m=6, x_M=10) + + expected = np.array([0., 2.5, 0.]) + (7. if increment else 0.) + np.testing.assert_allclose(rcv.data, expected) + + def test_zero(self): + grid = Grid(shape=(17,)) + rcv = SparseFunction(name='rcv', grid=grid, npoint=3) + rcv.data[:] = 7. + op = Operator(rcv.interpolate(0), name='ZeroSparseSum') + op.apply() + np.testing.assert_array_equal(rcv.data, 0.) + + def test_dtype(self): + grid = Grid(shape=(17,), dtype=np.float64) + f = Function(name='f', grid=grid) + rcv = SparseFunction(name='rcv', grid=grid, npoint=3, dtype=np.float32) + exprs = rcv.interpolate(f) + reduction = exprs.evaluate[-1].rhs + + assert reduction.dtype is rcv.dtype + + op = Operator(exprs, name='MixedPrecisionSparseSum', opt='noop') + values = [i for i in FindSymbols().visit(op) if isinstance(i, Temp)] + assert len(values) == 1 + assert values[0].dtype is reduction.dtype + + # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- diff --git a/tests/test_ir.py b/tests/test_ir.py index cd7bbee4b9..ab2685b44e 100644 --- a/tests/test_ir.py +++ b/tests/test_ir.py @@ -7,6 +7,7 @@ Constant, Dimension, Eq, Function, Grid, Inc, Operator, SubDimension, TimeFunction, switchconfig ) +from devito.finite_differences.differentiable import IndexSum, LocalSum from devito.ir.cgen import ccode from devito.ir.clusters import Cluster, ClusterGroup from devito.ir.equations import LoweredEq @@ -22,11 +23,12 @@ Any, Backward, Forward, Interval, IntervalGroup, IterationInterval, IterationSpace, NullInterval, null_ispace ) -from devito.symbolics import DefFunction, FieldFromPointer +from devito.symbolics import DefFunction, FieldFromPointer, uxreplace from devito.tools import prod from devito.tools.data_structures import frozendict from devito.types import ( - Array, Bundle, CriticalRegion, CustomDimension, Jump, Scalar, Symbol + Array, Bundle, ConditionalDimension, CriticalRegion, CustomDimension, Jump, Scalar, + Symbol ) @@ -1188,6 +1190,21 @@ def test_dimension_sort(self, expr, expected): assert list(dimension_sort(expr)) == eval(expected) + def test_reduction_dimensions(self): + r = CustomDimension('r', 0, 2, 3) + f = Function(name='f', dimensions=(r,), shape=(3,)) + acc = Symbol(name='acc', dtype=f.dtype) + + # A reduction loop can be requested without an index in the expression + expr = Inc(acc, 1, implicit_dims=r) + assert r not in expr.free_symbols + assert LoweredEq(expr).ispace.itdims == (r,) + + # A symbolic sum owns its loop; a separate free use still needs an outer loop + reduction = IndexSum(f[r], r) + assert LoweredEq(Eq(acc, reduction)).ispace.itdims == () + assert LoweredEq(Eq(acc, reduction + f[r])).ispace.itdims == (r,) + class TestCluster: @@ -1227,6 +1244,36 @@ def test_eq_hash_include_ispace(self): assert cgroup0 != cgroup1 assert len({cgroup0, cgroup1}) == 2 + def test_local_sums(self): + grid = Grid(shape=(4,)) + x, = grid.dimensions + f = Function(name='f', grid=grid) + i = ConditionalDimension('i', CustomDimension('i', 0, 1, 2), + condition=S.true, indirect=True) + j = ConditionalDimension('j', CustomDimension('j', 0, 1, 2), + condition=S.true, indirect=True) + inner = LocalSum(f[x + i], cdims=(i,)) + outer = LocalSum(inner*f[x + j], cdims=(j,)) + expr = LoweredEq(Eq(f[x], outer)) + assert expr.ispace.itdims == (x,) + cluster = Cluster(expr, expr.ispace) + group = ClusterGroup([cluster, cluster]) + + # Nested sums must lower first; occurrences in separate equations must + # remain distinct, since intervening writes may change the summand + assert expr.local_sums == cluster.local_sums == (inner, outer) + assert group.local_sums == (inner, outer, inner, outer) + assert cluster.local_sums is cluster.local_sums + assert group.local_sums is group.local_sums + + # Rebuilding expressions and clusters must not retain stale cached sums + replaced = cluster.exprs[0].apply(lambda e: uxreplace(e, {outer: inner})) + rebuilt = cluster.rebuild(exprs=[replaced]) + assert rebuilt.local_sums == (inner,) + assert ClusterGroup([rebuilt]).local_sums == (inner,) + assert cluster.local_sums == (inner, outer) + assert not cluster.rebuild(exprs=[Eq(f[x], 0)]).local_sums + class TestGuards: diff --git a/tests/test_pickle.py b/tests/test_pickle.py index 3f8a122bcc..954d39cc80 100644 --- a/tests/test_pickle.py +++ b/tests/test_pickle.py @@ -9,7 +9,7 @@ from devito import ( MPI, ConditionalDimension, Constant, Dimension, Eq, Function, Grid, IncrDimension, Min, Operator, PrecomputedSparseTimeFunction, SparseFunction, SteppingDimension, - SubDimension, SubDomain, TimeDimension, TimeFunction, solve + SubDimension, SubDomain, TimeDimension, TimeFunction, floor, solve ) from devito.data import LEFT, OWNED from devito.finite_differences.tools import centered, direct, left, right, transpose @@ -20,7 +20,7 @@ ) from devito.symbolics import ( CallFromPointer, Cast, DefFunction, FieldFromPointer, IntDiv, ListInitializer, SizeOf, - pow_to_mul + indexify, pow_to_mul ) from devito.tools import EnrichedTuple from devito.types import ( @@ -671,6 +671,23 @@ def test_pow_to_mul(self, pickle): assert new_expr.is_Mul + def test_local_sum(self, pickle): + grid = Grid(shape=(17, 19), dtype=np.float64) + f = Function(name='f', grid=grid, space_order=8) + rcv = SparseFunction(name='rcv', grid=grid, npoint=3, + interpolation='sinc', r=4) + reduction = indexify(rcv.interpolate(f).evaluate[-1].rhs) + rebuilt = pickle.loads(pickle.dumps(reduction)) + + assert type(rebuilt) is type(reduction) + assert str(rebuilt.expr) == str(reduction.expr) + assert tuple((d.symbolic_min, d.symbolic_max) + for d in rebuilt.dimensions) == ((-3, 4), (-3, 4)) + assert set(rebuilt.cdims) <= rebuilt.expr.free_symbols + assert tuple(map(str, rebuilt.conditionals.values())) == \ + tuple(map(str, reduction.conditionals.values())) + assert rebuilt.free_symbols.isdisjoint(rebuilt.bound_symbols) + class TestAdvanced: @@ -830,21 +847,12 @@ def test_collected_coeffs(self, pickle): def test_elemental(self, pickle): """ - Tests that elemental functions don't get reconstructed differently. + Tests that elementary functions don't get reconstructed differently. """ - grid = Grid(shape=(101, 101)) - time_range = TimeAxis(start=0.0, stop=1000.0, num=12) - - nrec = 101 - rec = Receiver(name='rec', grid=grid, npoint=nrec, time_range=time_range) - - u = TimeFunction(name="u", grid=grid, time_order=2, space_order=2) - rec_term = rec.interpolate(expr=u) - - eq = rec_term.evaluate[3] - eq = eq.func(eq.lhs, eq.rhs.args[0]) + grid = Grid(shape=(4, 4)) + f = Function(name='f', grid=grid) - op = Operator(eq) + op = Operator(Eq(f, floor(f + 0.5))) pkl_op = pickle.dumps(op) new_op = pickle.loads(pkl_op) diff --git a/tests/test_symbolics.py b/tests/test_symbolics.py index d9940d6e3e..552a7542c9 100644 --- a/tests/test_symbolics.py +++ b/tests/test_symbolics.py @@ -904,6 +904,7 @@ class TestUxreplace: @pytest.mark.parametrize('expr,subs,expected', [ ('f', '{f: g}', 'g'), + ('x + y', '{x: 0}', 'y'), ('f[x, y+1]', '{f.indexed: g.indexed}', 'g[x, y+1]'), ('cos(f)', '{cos: sin}', 'sin(f)'), ('cos(f + sin(g))', '{cos: sin, sin: cos}', 'sin(f + cos(g))'),