Skip to content

Commit 4d2e2b8

Browse files
committed
compiler: Introduce SparseLocalSum
1 parent cf24a2b commit 4d2e2b8

14 files changed

Lines changed: 318 additions & 36 deletions

File tree

devito/finite_differences/differentiable.py

Lines changed: 79 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,8 +21,8 @@
2121
from devito.finite_differences.tools import coeff_priority, make_shift_x0
2222
from devito.logger import warning
2323
from devito.tools import (
24-
Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict, infer_dtype,
25-
is_integer, is_number, memoized_func, split
24+
Pickable, Tag, as_tuple, extract_dtype, filter_ordered, flatten, frozendict,
25+
infer_dtype, is_integer, is_number, memoized_func, split
2626
)
2727
from devito.types import Array, DimensionTuple, Evaluable, StencilDimension
2828
from devito.types.basic import AbstractFunction, Indexed
@@ -36,6 +36,7 @@
3636
'IndexDerivative',
3737
'IndexDerivativeProperty',
3838
'Real',
39+
'SparseLocalSum',
3940
'Weights',
4041
]
4142

@@ -947,6 +948,82 @@ def free_symbols(self):
947948
func = DifferentiableOp._rebuild
948949

949950

951+
class SparseLocalSum(IndexSum, Pickable):
952+
953+
"""
954+
A zero-initialized sum over the guarded taps of a sparse interpolation.
955+
956+
`dimensions` are bound StencilDimensions; `predicates` contains their bounds
957+
conditions in the same order. Masked taps contribute zero.
958+
The sum remains symbolic until Cluster lowering chooses its implementation.
959+
960+
Examples
961+
--------
962+
For bilinear interpolation, `posx` and `posy` are the grid indices of sparse
963+
point `p`, and `wx` and `wy` hold its interpolation weights::
964+
965+
i = StencilDimension('i', 0, 1)
966+
j = StencilDimension('j', 0, 1)
967+
value = SparseLocalSum(
968+
wx[p, i]*wy[p, j]*f[posx + i, posy + j],
969+
dimensions=(i, j), dtype=f.dtype,
970+
predicates=(And(posx + i >= x_m, posx + i <= x_M),
971+
And(posy + j >= y_m, posy + j <= y_M))
972+
)
973+
Eq(rcv[p], value)
974+
975+
The scalar lowering has the following semantics (pseudocode)::
976+
977+
acc = 0
978+
for i in range(2):
979+
for j in range(2):
980+
if x_m <= posx + i <= x_M and y_m <= posy + j <= y_M:
981+
acc += wx[p, i]*wy[p, j]*f[posx + i, posy + j]
982+
rcv[p] = acc
983+
984+
`i` and `j` are local to the sum; `p` remains an outer iteration dimension.
985+
If every tap is masked, `rcv[p]` receives zero.
986+
"""
987+
988+
__rargs__ = ('expr',)
989+
__rkwargs__ = ('dimensions', 'dtype', 'predicates')
990+
991+
def __new__(cls, expr, dimensions, dtype, predicates, **kwargs):
992+
# Unlike IndexSum, the summand can be independent of a bound index, e.g. zero
993+
obj = sympy.Expr.__new__(cls, expr)
994+
obj._expr = expr
995+
obj._dimensions = as_tuple(dimensions)
996+
obj._dtype = np.dtype(dtype).type
997+
obj._predicates = as_tuple(predicates)
998+
return obj
999+
1000+
def _hashable_content(self):
1001+
return super()._hashable_content() + (self.dtype, self.predicates)
1002+
1003+
@property
1004+
def dtype(self):
1005+
return self._dtype
1006+
1007+
@property
1008+
def predicates(self):
1009+
return self._predicates
1010+
1011+
@property
1012+
def free_symbols(self):
1013+
symbols = super().free_symbols.union(*[c.free_symbols for c in self.predicates])
1014+
return symbols.difference(self.dimensions)
1015+
1016+
def _evaluate(self, **kwargs):
1017+
return self._rebuild(Evaluable._evaluate_maybe_nested(self.expr, **kwargs))
1018+
1019+
def _xreplace(self, rule):
1020+
from devito.symbolics import uxreplace
1021+
rebuilt = uxreplace(self, {k: sympy.sympify(v) for k, v in rule.items()})
1022+
return rebuilt, rebuilt is not self
1023+
1024+
__reduce_ex__ = Pickable.__reduce_ex__
1025+
1026+
9501027
class WeightsIndexed(Indexed):
9511028

9521029
@property

devito/ir/clusters/cluster.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -152,6 +152,11 @@ def dist_dimensions(self):
152152
def scope(self):
153153
return Scope(self.exprs)
154154

155+
@cached_property
156+
def sparse_sums(self):
157+
"""The sparse sums in equation order, retaining occurrences across equations."""
158+
return tuple(s for e in self.exprs for s in e.sparse_sums)
159+
155160
@cached_property
156161
def functions(self):
157162
return self.scope.functions
@@ -674,6 +679,10 @@ def rebuild(self, **kwargs):
674679
def exprs(self):
675680
return flatten(c.exprs for c in self)
676681

682+
@cached_property
683+
def sparse_sums(self):
684+
return tuple(s for c in self for s in c.sparse_sums)
685+
677686
@cached_property
678687
def scope(self):
679688
return Scope(exprs=self.exprs)

devito/ir/equations/algorithms.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,9 @@ def handle_indexed(indexed):
3333
for i in indexed.indices:
3434
try:
3535
# Assume it's an AffineIndexAccessFunction...
36-
relation.append(i.d)
36+
# It may contain only a scalar offset and bound stencil indices
37+
if i.d:
38+
relation.append(i.d)
3739
except AttributeError:
3840
# It's not! Maybe there are some nested Indexeds (e.g., the
3941
# situation is A[B[i]])

devito/ir/equations/equation.py

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,15 @@
44
import numpy as np
55
import sympy
66

7-
from devito.finite_differences.differentiable import diff2sympy
7+
from devito.finite_differences.differentiable import SparseLocalSum, diff2sympy
88
from devito.ir.equations.algorithms import dimension_sort, generate_conditionals
99
from devito.ir.support import (
1010
Interval, IntervalGroup, IterationSpace, Stencil, detect_accesses
1111
)
12-
from devito.symbolics import limits_mapper, retrieve_accesses
12+
from devito.symbolics import limits_mapper, retrieve_accesses, search
1313
from devito.tools import (
14-
Pickable, Tag, as_hashable, filter_sorted, frozendict, reuse_if_unchanged
14+
Pickable, Tag, as_hashable, filter_ordered, filter_sorted, frozendict,
15+
reuse_if_unchanged
1516
)
1617
from devito.types import Eq, Inc, ReduceMax, ReduceMin, ReduceMinMax
1718

@@ -50,6 +51,11 @@ def ispace(self):
5051
def dimensions(self):
5152
return set(self.ispace.dimensions)
5253

54+
@cached_property
55+
def sparse_sums(self):
56+
"""The sparse sums in dependency order, with nested sums first."""
57+
return tuple(filter_ordered(search(self, SparseLocalSum, mode='all')))
58+
5359
@property
5460
def implicit_dims(self):
5561
return self._implicit_dims

devito/operations/interpolators.py

Lines changed: 20 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -11,14 +11,16 @@
1111
except ImportError:
1212
from numpy import i0
1313

14-
from devito.finite_differences.differentiable import Mul
14+
from devito.finite_differences.differentiable import Mul, SparseLocalSum
1515
from devito.finite_differences.elementary import floor
1616
from devito.logger import warning
1717
from devito.symbolics import INT, retrieve_function_carriers, retrieve_functions
1818
from devito.tools import (
1919
Pickable, as_fp64_decimal, as_list, as_tuple, filter_ordered, flatten, memoized_meth
2020
)
21-
from devito.types import CustomDimension, Eq, Evaluable, Inc, SubFunction, Symbol
21+
from devito.types import (
22+
CustomDimension, Eq, Evaluable, Inc, StencilDimension, SubFunction
23+
)
2224
from devito.types.utils import DimensionTuple
2325

2426
__all__ = ['LinearInterpolator', 'NearestInterpolator',
@@ -453,19 +455,21 @@ def _interp_idx(self, variables, implicit_dims=None, subdomain=None,
453455

454456
return idx_subs, temps
455457

456-
def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None):
458+
def _local_accumulator(self, expr, idx_subs, subdomain=None):
457459
"""
458-
Generate a local accumulator for the interpolation/injection operation.
460+
Represent the local sum of weighted interpolation contributions.
459461
"""
460-
# Accumulate point-wise contributions into a temporary
461-
rhs = Symbol(name=f'sum{self.sfunction.name}', dtype=self.sfunction.dtype)
462-
summands = [Eq(rhs, 0., implicit_dims=implicit_dims)]
463-
# Substitute coordinate base symbols into the interpolation coefficients
464462
weights = self._weights(subdomain=subdomain)
465-
summands.extend([Inc(rhs, (weights * expr).xreplace(idx_subs),
466-
implicit_dims=implicit_dims)])
467-
468-
return summands, rhs
463+
rdims = self._rdim(subdomain=subdomain)
464+
# Bound stencil indices are independent of the outer sparse-point dimension
465+
dimensions = tuple(StencilDimension(d.name, d.parent.symbolic_min,
466+
d.parent.symbolic_max) for d in rdims)
467+
aliases = {a: i for d, i in zip(rdims, dimensions, strict=True)
468+
for a in (d, d.parent)}
469+
summand = (weights * expr).xreplace(idx_subs).xreplace(aliases)
470+
predicates = tuple(d.condition.xreplace(aliases) for d in rdims)
471+
return SparseLocalSum(summand, dimensions, self.sfunction.dtype,
472+
predicates=predicates)
469473

470474
@check_radius
471475
@check_coords
@@ -539,16 +543,13 @@ def _interpolate(self, expr, increment=False, self_subs=None, implicit_dims=None
539543
idx_subs, temps = self._interp_idx(variables, implicit_dims=implicit_dims,
540544
subdomain=subdomain)
541545

542-
# Local scalar for accumulation over radius
543-
summands, rhs = self._local_accumulator(expr, idx_subs,
544-
implicit_dims=implicit_dims,
545-
subdomain=subdomain)
546+
rhs = self._local_accumulator(expr, idx_subs, subdomain=subdomain)
546547
# Write/Incr `self`
547548
lhs = self.sfunction.subs(self_subs)
548549
ecls = Inc if increment else Eq
549550
last = [ecls(lhs, rhs, implicit_dims=implicit_dims)]
550551

551-
return temps + summands + last
552+
return temps + last
552553

553554
def _inject(self, field, expr, increment=True, implicit_dims=None):
554555
"""
@@ -785,8 +786,8 @@ class NearestInterpolator(LinearInterpolator):
785786

786787
_name = 'nearest'
787788

788-
def _local_accumulator(self, expr, idx_subs, implicit_dims=None, subdomain=None):
789-
return [], expr.xreplace(idx_subs)
789+
def _local_accumulator(self, expr, idx_subs, subdomain=None):
790+
return expr.xreplace(idx_subs)
790791

791792
@memoized_meth
792793
def _rdim(self, subdomain=None, shifts=None):

devito/operator/operator.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,8 @@
3333
from devito.parameters import configuration
3434
from devito.passes import (
3535
Graph, error_mapper, finalize_args, generate_implicit, generate_macros, is_on_device,
36-
lower_dtypes, lower_index_derivatives, minimize_symbols, optimize_pows, unevaluate
36+
lower_dtypes, lower_index_derivatives, lower_sparse_sums, minimize_symbols,
37+
optimize_pows, unevaluate
3738
)
3839
from devito.symbolics import estimate_cost, subs_op_args
3940
from devito.tools import (
@@ -423,6 +424,7 @@ def _lower_clusters(cls, expressions, profiler=None, **kwargs):
423424
clusters = generate_implicit(clusters)
424425

425426
# Lower all remaining high order symbolic objects
427+
clusters = lower_sparse_sums(clusters, **kwargs)
426428
clusters = lower_index_derivatives(clusters, **kwargs)
427429

428430
# Turn pows into multiplications. This must happen as late as possible

devito/passes/clusters/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,4 +9,5 @@
99
from .implicit import * # noqa
1010
from .misc import * # noqa
1111
from .derivatives import * # noqa
12+
from .sparse import * # noqa
1213
from .unevaluate import * # noqa

devito/passes/clusters/cse.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
# Moved in 1.13
1212
from sympy.core.basic import ordering_of_classes
1313

14-
from devito.finite_differences.differentiable import IndexDerivative
14+
from devito.finite_differences.differentiable import IndexSum
1515
from devito.ir import Cluster, Scope, cluster_pass
1616
from devito.symbolics import (
1717
DefFunction, Reserved, estimate_cost, q_leaf, q_terminal, search
@@ -427,7 +427,7 @@ def _(expr):
427427
return {}
428428

429429

430-
@_catch.register(IndexDerivative)
430+
@_catch.register(IndexSum)
431431
def _(expr):
432432
"""
433433
Handler for symbol-binding objects. There can be many of them and therefore

devito/passes/clusters/sparse.py

Lines changed: 82 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,82 @@
1+
from functools import partial
2+
3+
from devito.ir import ClusterizedEq, Interval, IterationSpace
4+
from devito.symbolics import uxreplace
5+
from devito.tools import timed_pass
6+
from devito.types import ConditionalDimension, Eq, Inc, Temp
7+
8+
__all__ = ['lower_sparse_sum', 'lower_sparse_sums']
9+
10+
11+
@timed_pass()
12+
def lower_sparse_sums(clusters, sregistry=None, **kwargs):
13+
"""
14+
Lower SparseLocalSums into a private initializer and guarded accumulation.
15+
16+
For example, with `i = StencilDimension('i', 0, 1)`::
17+
18+
Eq(rcv[p], SparseLocalSum(
19+
w[p, i]*f[pos + i], dimensions=(i,), dtype=f.dtype,
20+
predicates=(And(pos + i >= x_m, pos + i <= x_M),)
21+
))
22+
23+
becomes the following computation at each sparse point `p` (pseudocode)::
24+
25+
sum0 = 0
26+
for i in range(2):
27+
if x_m <= pos + i <= x_M:
28+
sum0 += w[p, i]*f[pos + i]
29+
rcv[p] = sum0
30+
31+
The initializer and result stay outside the tap guard, so even a fully
32+
masked stencil assigns zero to `rcv[p]`.
33+
"""
34+
processed = []
35+
for c in clusters:
36+
if not c.sparse_sums:
37+
processed.append(c)
38+
continue
39+
40+
for e in c.exprs:
41+
subs = {}
42+
for reduction in e.sparse_sums:
43+
init, update, value = lower_sparse_sum(
44+
c, uxreplace(reduction, subs), sregistry
45+
)
46+
processed.extend([init, update])
47+
subs[reduction] = value
48+
49+
expr = e.apply(partial(uxreplace, rule=subs))
50+
processed.append(c.rebuild(exprs=[expr]))
51+
52+
return processed
53+
54+
55+
def lower_sparse_sum(cluster, reduction, sregistry):
56+
"""
57+
Construct the private initializer and guarded accumulation for one sum.
58+
"""
59+
value = Temp(name=sregistry.make_name(prefix='sum'), dtype=reduction.dtype)
60+
61+
dims = reduction.dimensions
62+
inner = IterationSpace([Interval(d) for d in dims])
63+
ispace = IterationSpace.union(
64+
cluster.ispace, inner, relations=(cluster.ispace.itdims + dims,)
65+
)
66+
67+
init = cluster.rebuild(exprs=[Eq(value, 0)])
68+
conditionals = {
69+
ConditionalDimension(d.name, parent=d, condition=v, indirect=True): v
70+
for d, v in zip(dims, reduction.predicates, strict=True)
71+
}
72+
# Cluster construction does not infer conditionals from a plain Eq/Inc,
73+
# so attach the per-tap guards explicitly
74+
expr = ClusterizedEq(
75+
Inc(value, reduction.expr), ispace=ispace, conditionals=conditionals
76+
)
77+
78+
properties = cluster.properties.sequentialize(dims)
79+
80+
update = cluster.rebuild(exprs=expr, ispace=ispace, properties=properties)
81+
82+
return init, update, value

devito/symbolics/inspection.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from sympy.core.numbers import ImaginaryUnit
88

99
from devito.finite_differences import Derivative
10-
from devito.finite_differences.differentiable import IndexDerivative
10+
from devito.finite_differences.differentiable import IndexDerivative, IndexSum
1111
from devito.logger import warning
1212
from devito.symbolics.extended_dtypes import INT
1313
from devito.symbolics.extended_sympy import CallFromPointer, Cast, DefFunction, Reserved
@@ -265,7 +265,7 @@ def _(expr, estimate, seen):
265265
return _estimate_cost(expr._evaluate(expand=False), estimate, seen)
266266

267267

268-
@_estimate_cost.register(IndexDerivative)
268+
@_estimate_cost.register(IndexSum)
269269
@dont_count_if_seen
270270
def _(expr, estimate, seen):
271271
flops, _ = _estimate_cost(expr.expr, estimate, seen)
@@ -275,7 +275,7 @@ def _(expr, estimate, seen):
275275

276276
# To be multiplied by the number of points this index sum implicitly
277277
# iterates over
278-
flops *= prod(i._size for i in expr.dimensions)
278+
flops *= prod(i.symbolic_size for i in expr.dimensions)
279279

280280
return flops, False
281281

0 commit comments

Comments
 (0)