Skip to content

Commit 531cd55

Browse files
committed
compiler: Tweak LocalSum
1 parent c3b939a commit 531cd55

4 files changed

Lines changed: 29 additions & 9 deletions

File tree

‎devito/finite_differences/differentiable.py‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -994,23 +994,30 @@ class LocalSum(IndexSum, Pickable):
994994
"""
995995

996996
__rargs__ = ('expr',)
997-
__rkwargs__ = ('cdims',)
997+
__rkwargs__ = ('cdims', 'dtype')
998998

999-
def __new__(cls, expr, cdims=(), **kwargs):
999+
def __new__(cls, expr, cdims=(), dtype=None, **kwargs):
10001000
obj = sympy.Expr.__new__(cls, expr)
10011001

10021002
obj._expr = expr
10031003
obj._cdims = as_tuple(cdims)
1004+
obj._dtype = dtype
10041005

10051006
return obj
10061007

10071008
def _hashable_content(self):
1008-
return super()._hashable_content() + (self.cdims,)
1009+
return super()._hashable_content() + (self.cdims, self.dtype)
10091010

10101011
@property
10111012
def cdims(self):
10121013
return self._cdims
10131014

1015+
@cached_property
1016+
def dtype(self):
1017+
if self._dtype is None:
1018+
return extract_dtype(self.expr)
1019+
return self._dtype
1020+
10141021
@cached_property
10151022
def dimensions(self):
10161023
return tuple(d.parent for d in self.cdims)

‎devito/operations/interpolators.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -460,7 +460,7 @@ def _local_accumulator(self, expr, idx_subs, subdomain=None):
460460
weights = self._weights(subdomain=subdomain)
461461
rdims = self._rdim(subdomain=subdomain)
462462
summand = (weights * expr).xreplace(idx_subs)
463-
return LocalSum(summand, cdims=rdims)
463+
return LocalSum(summand, cdims=rdims, dtype=self.sfunction.dtype)
464464

465465
@check_radius
466466
@check_coords

‎devito/passes/clusters/localsum.py‎

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,3 @@
1-
from functools import partial
2-
31
from devito.ir import ClusterizedEq, Interval, IterationSpace
42
from devito.symbolics import uxreplace
53
from devito.tools import timed_pass
@@ -46,7 +44,7 @@ def lower_local_sums(clusters, sregistry=None, **kwargs):
4644
processed.extend([init, update])
4745
subs[reduction] = value
4846

49-
expr = e.apply(partial(uxreplace, rule=subs))
47+
expr = uxreplace(e, subs)
5048
processed.append(c.rebuild(exprs=[expr]))
5149

5250
return processed
@@ -56,7 +54,7 @@ def lower_local_sum(cluster, reduction, sregistry):
5654
"""
5755
Construct the private initializer and guarded accumulation for one sum.
5856
"""
59-
value = Temp(name=sregistry.make_name(prefix='sum'), dtype=cluster.dtype)
57+
value = Temp(name=sregistry.make_name(prefix='sum'), dtype=reduction.dtype)
6058

6159
dims = reduction.dimensions
6260
inner = IterationSpace([Interval(d) for d in dims])

‎tests/test_interpolation.py‎

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,12 +12,13 @@
1212
switchconfig
1313
)
1414
from devito.finite_differences import LocalSum
15-
from devito.ir import LoweredEq
15+
from devito.ir import FindSymbols, LoweredEq
1616
from devito.operations.interpolators import (
1717
LinearInterpolator, SincInterpolator, _cell_indices
1818
)
1919
from devito.symbolics import uxreplace
2020
from devito.tools import as_tuple
21+
from devito.types import Temp
2122
from examples.seismic import (
2223
AcquisitionGeometry, Receiver, RickerSource, TimeAxis, demo_model
2324
)
@@ -85,6 +86,20 @@ def test_zero(self):
8586
op.apply()
8687
np.testing.assert_array_equal(rcv.data, 0.)
8788

89+
def test_dtype(self):
90+
grid = Grid(shape=(17,), dtype=np.float64)
91+
f = Function(name='f', grid=grid)
92+
rcv = SparseFunction(name='rcv', grid=grid, npoint=3, dtype=np.float32)
93+
exprs = rcv.interpolate(f)
94+
reduction = exprs.evaluate[-1].rhs
95+
96+
assert reduction.dtype is rcv.dtype
97+
98+
op = Operator(exprs, name='MixedPrecisionSparseSum', opt='noop')
99+
values = [i for i in FindSymbols().visit(op) if isinstance(i, Temp)]
100+
assert len(values) == 1
101+
assert values[0].dtype is reduction.dtype
102+
88103

89104
# ---------------------------------------------------------------------------
90105
# Helpers

0 commit comments

Comments
 (0)