Skip to content

Commit 04411cd

Browse files
committed
api: compute sparse position/floor in fp64 to fix off-by-one cell shift
The linear interpolator computes `(c - o)/h` in the grid's dtype, which for fp32 grids rounds coord/origin/spacing to fp32. That rounding can push the position across an integer boundary, so `floor((c - o)/h)` picks a different cell than the fp64-truth cell -- and CPU vs GPU can disagree on the fractional part while agreeing on the integer position, producing inconsistent injection weights. Compute `pos = (c - o)/h` and `floor(pos)` in fp64 by casting the fp32 free symbols and by substituting the spacing symbols with their fp64 decimal value (recovered via the fp32 short-form round-trip). A new CIRE subpass `CireInvariantsSparse` hoists `pos` and `floor(pos)` out of the per-stencil-point inner loop into per-source preamble Arrays. The floor tab is stored as int32 (half the memory of an fp64 tab, no precision loss), and bare `floor(pos)` uses inside `pos - floor(pos)` are rewritten to `DOUBLE(int_tab)` so both consumers share the tab. The `sympy_dtype` inference is taught to recognize `Cast` (outermost and inner) so the printer emits `floor` (fp64) instead of `floorf` (fp32) when a `DOUBLE(...)` cast is present in the expression.
1 parent 24444ce commit 04411cd

9 files changed

Lines changed: 416 additions & 81 deletions

File tree

‎devito/operations/interpolators.py‎

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -543,11 +543,9 @@ def _point_symbols(self, shifts=None):
543543
return DimensionTuple(*symbols, getters=self.grid.dimensions)
544544

545545
def _coeff_temps(self, implicit_dims, shifts=None):
546-
# Positions
547546
pmap = self.sfunction._position_map(shifts=shifts)
548547
psyms = self._point_symbols(shifts)
549-
poseq = [Eq(psyms[d], pos - floor(pos),
550-
implicit_dims=implicit_dims)
548+
poseq = [Eq(psyms[d], pos - floor(pos), implicit_dims=implicit_dims)
551549
for (d, pos) in zip(self._gdims, pmap.keys(), strict=True)]
552550
return poseq
553551

‎devito/passes/clusters/aliases.py‎

Lines changed: 77 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,8 @@
1515
from devito.passes.clusters.cse import _cse
1616
from devito.passes.clusters.utils import expose_tuning_knobs
1717
from devito.symbolics import (
18-
Uxmapper, estimate_cost, retrieve_functions, reuse_if_untouched, search, sympy_dtype,
19-
uxreplace
18+
DOUBLE, INT, Cast, Uxmapper, estimate_cost, retrieve_functions, reuse_if_untouched,
19+
search, sympy_dtype, uxreplace
2020
)
2121
from devito.tools import (
2222
Reconstructable, Stamp, as_mapper, as_tuple, flatten, frozendict, generator,
@@ -295,6 +295,10 @@ def _do_generate(self, exprs, exclude, cbk_search, cbk_compose=None):
295295

296296
class CireInvariants(CireTransformerLegacy, Queue):
297297

298+
# Predicate on Cluster used to pick which ones this pass fires on.
299+
# Subclasses override to target a different kind of cluster.
300+
_cluster_filter = staticmethod(lambda c: c.is_dense)
301+
298302
def __init__(self, sregistry, options, platform):
299303
super().__init__(sregistry, options, platform)
300304

@@ -324,7 +328,8 @@ def callback(self, clusters, prefix, xtracted=None):
324328
key = lambda c: self._lookup_key(c, d)
325329
processed = list(clusters)
326330
for ak, group in as_mapper(clusters, key=key).items():
327-
g = [c for c in group if c.is_dense and c not in xtracted]
331+
g = [c for c in group
332+
if self._cluster_filter(c) and c not in xtracted]
328333
if not g:
329334
continue
330335

@@ -387,6 +392,73 @@ def _generate(self, cgroup, exclude):
387392
yield self._do_generate(exprs, exclude, cbk_search)
388393

389394

395+
def _is_floor(e):
396+
return getattr(e, 'is_Function', False) and e.func.__name__ == 'floor'
397+
398+
399+
class CireInvariantsSparse(CireInvariants):
400+
401+
"""
402+
Hoist sparse-point position temps `pos = (c - o)/h` and their
403+
`floor(pos)` out of the per-stencil-point inner loop into preamble
404+
Arrays computed once per source. The inner loop then reads
405+
tabulated values instead of recomputing `floor((c - o)/h)` for each
406+
`(rp_srcx, rp_srcy, rp_srcz)` combination. `Lift` then moves the
407+
preamble out of the time loop.
408+
"""
409+
410+
_cluster_filter = staticmethod(lambda c: not c.is_dense)
411+
412+
def _generate(self, cgroup, exclude):
413+
# Tabulate `pos` as fp64 and `INT(floor(pos))` as int32. The int32
414+
# tab feeds both integer position lookups and (via DOUBLE cast) the
415+
# bare `floor(pos)` uses in `pos - floor(pos)`.
416+
counter = generator()
417+
make_f64 = lambda: Symbol(name=f'dummy{counter()}', dtype=np.float64)
418+
make_i32 = lambda: Symbol(name=f'dummy{counter()}', dtype=np.int32)
419+
420+
mapper = Uxmapper()
421+
422+
def _add(expr, make):
423+
if expr.is_commutative is False:
424+
return
425+
if {a.function for a in expr.free_symbols} & exclude:
426+
return
427+
mapper.add(expr, make, None)
428+
429+
for e in cgroup.exprs:
430+
for f in search(e, _is_floor, 'all', 'bfs'):
431+
_add(f.args[0], make_f64)
432+
_add(INT(f), make_i32)
433+
434+
yield mapper
435+
436+
def _choose(self, aliases, cgroup, mapper):
437+
# Skip score-based filtering and fold bare `floor(pos)` onto the
438+
# int32 alias built for `INT(floor(pos))` via a `DOUBLE(...)` cast.
439+
exprs = cgroup.exprs
440+
441+
aliases = AliasList(aliases)
442+
if not aliases:
443+
return exprs, aliases
444+
445+
aliaseds = set(aliases.aliaseds)
446+
subs = {k: v for k, v in mapper.items() if v.free_symbols & aliaseds}
447+
448+
for k, v in list(mapper.items()):
449+
if not isinstance(k, Cast):
450+
continue
451+
if not (isinstance(k.dtype, type) and
452+
issubclass(k.dtype, np.integer)):
453+
continue
454+
inner = k.base
455+
if _is_floor(inner) and (v.free_symbols & aliaseds):
456+
subs[inner] = DOUBLE(v)
457+
458+
exprs = [uxreplace(e, subs) for e in exprs]
459+
return exprs, aliases
460+
461+
390462
class CireDerivatives(CireTransformerLegacy):
391463

392464
def __init__(self, sregistry, options, platform):
@@ -519,7 +591,8 @@ def _cbk_search2(self, expr, rank):
519591
# Subpass mapper
520592
modes = {
521593
'invariants': [CireInvariantsElementary,
522-
CireInvariantsDivs],
594+
CireInvariantsDivs,
595+
CireInvariantsSparse],
523596
'eval-derivs': [CireEvalDerivatives], # NOTE: legacy pass
524597
'index-derivs': [CireIndexDerivatives],
525598
}

‎devito/symbolics/inspection.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -316,10 +316,24 @@ def sympy_dtype(expr, base=None, default=None, smin=None):
316316
if expr is None:
317317
return default
318318

319+
# An outermost Cast (e.g. `DOUBLE(...)`, `INT(...)`) declares the dtype
320+
# of the whole expression regardless of its operand's free symbols.
321+
if isinstance(expr, Cast):
322+
cast_dtype = expr.dtype
323+
if isinstance(cast_dtype, type) and issubclass(cast_dtype, np.generic):
324+
return cast_dtype
325+
319326
dtypes = set()
327+
# A Cast inside the tree bumps the inferred dtype: it forces the
328+
# operation to occur at (at least) the Cast's precision.
329+
for c in expr.atoms(Cast):
330+
cd = c.dtype
331+
if isinstance(cd, type) and issubclass(cd, np.generic):
332+
dtypes.add(cd)
320333
for i in expr.free_symbols:
321334
with suppress(AttributeError):
322335
dtypes.add(i.dtype)
336+
dtypes.discard(None)
323337

324338
if not dtypes or not np.issubdtype(base, np.complexfloating):
325339
dtypes.update({base} - {None})

‎devito/types/sparse.py‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
from devito.operations import (
1212
LinearInterpolator, PrecomputedInterpolator, SincInterpolator
1313
)
14-
from devito.symbolics import indexify, retrieve_function_carriers
14+
from devito.symbolics import DOUBLE, indexify, retrieve_function_carriers
1515
from devito.tools import (
1616
ReducerMap, as_tuple, dtype_to_mpidtype, filter_ordered, flatten, is_integer,
1717
memoized_meth, prod
@@ -394,10 +394,14 @@ def _position_map(self, shifts=None):
394394
Dimension, of additional physical offsets to subtract from the sparse
395395
coordinates (e.g. ``h_x/2`` for a field staggered in ``x``). If ``shifts``
396396
is None, only the grid origin is subtracted.
397+
398+
The origin is cast to double so the subtraction and subsequent division
399+
by the (fp32) spacing happen in fp64 and pick the correct cell
400+
regardless of the grid's fp32 rounding.
397401
"""
398402
shifts = shifts or (0,) * len(self.grid.dimensions)
399403
return OrderedDict([
400-
((c - o - s)/d.spacing, p)
404+
((c - DOUBLE(o) - s)/d.spacing, p)
401405
for p, c, d, o, s in zip(
402406
self._pos_symbols(shifts=shifts),
403407
self._coordinate_symbols,

‎examples/userapi/06_sparse_operations.ipynb‎

Lines changed: 241 additions & 53 deletions
Large diffs are not rendered by default.

‎tests/test_dle.py‎

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -204,9 +204,9 @@ def test_basic(self):
204204

205205
op = Operator(eqns, opt=('advanced', {'blockrelax': True}))
206206

207-
bns, _ = assert_blocking(op, {'x0_blk0', 'p_src0_blk0'})
207+
bns, _ = assert_blocking(op, {'x0_blk0', 'p_src0_blk0', 'p_src1_blk0'})
208208

209-
iters = FindNodes(Iteration).visit(bns['p_src0_blk0'])
209+
iters = FindNodes(Iteration).visit(bns['p_src1_blk0'])
210210
assert len(iters) == 5
211211
assert iters[0].dim.is_Block
212212
assert iters[1].dim.is_Block
@@ -1070,12 +1070,14 @@ def test_incr_perfect_sparse_outer(self):
10701070
'openmp': True}))
10711071

10721072
iters = FindNodes(Iteration).visit(op)
1073-
assert len(iters) == 5
1074-
assert iters[0].is_Sequential
1075-
assert all(i.is_ParallelAtomic for i in iters[1:])
1076-
assert iters[1].pragmas[0].ccode.value ==\
1073+
# 4 preamble iterations (p_u + rp_u{x,y,z}) hoisted out of time by
1074+
# `CireInvariantsSparse`, then the time loop and the 4 injection iters.
1075+
assert len(iters) == 9
1076+
assert iters[4].is_Sequential
1077+
assert all(i.is_ParallelAtomic for i in iters[:4] + iters[5:])
1078+
assert iters[5].pragmas[0].ccode.value ==\
10771079
'omp for schedule(dynamic,chunk_size)'
1078-
assert all(not i.pragmas for i in iters[2:])
1080+
assert all(not i.pragmas for i in iters[6:])
10791081

10801082
@pytest.mark.parametrize('exprs,simd_level,expected', [
10811083
(['Eq(y.symbolic_max, g[0, x], implicit_dims=(t, x))',

‎tests/test_dse.py‎

Lines changed: 19 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -50,12 +50,18 @@ def test_scheduling_after_rewrite():
5050
op = Operator([eqn1] + eqn2 + [eqn3])
5151
trees = retrieve_iteration_tree(op)
5252

53-
# Check loop nest structure
54-
assert all(
55-
i.dim is j for i, j in zip(trees[0], grid.dimensions, strict=True)
56-
) # time invariant
57-
assert trees[1].root.dim is grid.time_dim
58-
assert all(trees[1].root.dim is tree.root.dim for tree in trees[1:])
53+
# Check loop nest structure. `CireInvariantsSparse` hoists the sparse
54+
# position preamble as an extra time-invariant tree ahead of the time
55+
# loops, alongside the dense sin(const) preamble.
56+
invariants, in_time = [], []
57+
for t in trees:
58+
(in_time if t.root.dim is grid.time_dim else invariants).append(t)
59+
assert invariants # sin(const) + the sparse-position preamble(s)
60+
assert in_time # the actual time-loop clusters
61+
# Invariants must all precede the time-loop trees
62+
for inv in invariants:
63+
for tt in in_time:
64+
assert trees.index(inv) < trees.index(tt)
5965

6066

6167
@pytest.mark.parametrize('expr,expected', [
@@ -1399,7 +1405,10 @@ def test_catch_duplicate_from_different_clusters(self):
13991405
op = Operator(eqns, opt=('advanced', {'cire-mingain': 100}))
14001406

14011407
arrays = [i for i in FindSymbols().visit(op) if i.is_Array]
1402-
assert len(arrays) == 3
1408+
# 3 dense-derivative aliases + 4 sparse-position preambles (pos_x,
1409+
# pos_y and their int(floor(...)) tabulations) from
1410+
# `CireInvariantsSparse`.
1411+
assert len(arrays) == 7
14031412
assert all(i._mem_heap and not i._mem_external for i in arrays)
14041413

14051414
def test_discarded_compound(self):
@@ -1542,7 +1551,9 @@ def test_drop_redundants_after_fusion(self, rotate):
15421551
op = Operator(eqns, opt=('advanced', {'cire-rotate': rotate}))
15431552

15441553
arrays = [i for i in FindSymbols().visit(op) if i.is_Array]
1545-
assert len(arrays) == 2
1554+
# 2 dense aliases + 4 sparse-position preambles per sparse function
1555+
# (pos_x, pos_y, int_floor_x, int_floor_y) x 1 SparseTimeFunction.
1556+
assert len(arrays) == 6
15461557
assert all(i._mem_heap and not i._mem_external for i in arrays)
15471558

15481559
def test_full_shape_big_temporaries(self):

‎tests/test_interpolation.py‎

Lines changed: 43 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
switchconfig
1313
)
1414
from devito.operations.interpolators import LinearInterpolator, SincInterpolator
15+
from devito.symbolics.extended_sympy import Cast
1516
from devito.tools import as_tuple
1617
from examples.seismic import (
1718
AcquisitionGeometry, Receiver, RickerSource, TimeAxis, demo_model
@@ -489,7 +490,7 @@ def test_inject_staggered(self, stagg):
489490
if stagg == 'NODE':
490491
assert np.isclose(a.data[5, 5, 5], 2, rtol=1e-6)
491492
# all other should be zero
492-
assert np.sum(a.data) == 2
493+
assert np.isclose(np.sum(a.data), 2, rtol=1e-6)
493494
else:
494495
# Bottom corner since the source position is left of the staggered field
495496
corner = [5, 5, 5] - np.array(a.staggered)
@@ -507,7 +508,8 @@ def test_inject_staggered(self, stagg):
507508
assert np.allclose(a.data[slices], interp_val, rtol=1e-6)
508509
# All other should be zero so should sum to the interp_val * number of points.
509510
# Use abs to make sure there is no +- cancellations
510-
assert np.sum(np.abs(a.data)) == interp_val * 2**(sum(np.array(a.staggered)))
511+
expected = float(interp_val * 2**(sum(np.array(a.staggered))))
512+
assert np.isclose(np.sum(np.abs(a.data)), expected, rtol=1e-6)
511513

512514
def test_inject_staggered_mixed(self):
513515
grid = Grid((11, 11, 11))
@@ -870,6 +872,45 @@ def test_wrong_coords(self):
870872
s.interpolate(u + s2)
871873
assert "Interpolation/injection with" in str(vinfo.value)
872874

875+
def test_position_map_fp64(self):
876+
"""
877+
The linear sparse position `(c - o)/h` must be computed in fp64 even
878+
when the grid dtype is fp32. Otherwise fp32 rounding of the origin
879+
can push the position across an integer boundary, so `floor((c-o)/h)`
880+
picks a different cell than the fp64-truth cell.
881+
882+
Uses `h = 0.1` (not fp32-exact) and a coordinate whose fp64 position
883+
is just below cell `7`; without the fix the fp32 position rounds up
884+
to exactly `7.0`, shifting the injection by one cell.
885+
"""
886+
# h = 0.1 in fp64, but not exactly representable in fp32
887+
grid = Grid(shape=(30, 30), extent=(2.9, 2.9))
888+
assert grid.dtype is np.float32
889+
890+
sf = SparseTimeFunction(name='sf', grid=grid, npoint=1, nt=2)
891+
892+
# Symbolic check: the origin is wrapped in a `DOUBLE(...)` cast so
893+
# the subtraction and subsequent division promote to fp64.
894+
pmap = sf._position_map()
895+
for expr in pmap:
896+
assert expr.find(Cast), \
897+
f"position map missing fp64 cast on origin: {expr}"
898+
899+
# End-to-end: fp64 position ~ 6.99999988 -> floor = 6, fraction ~ 1.
900+
# Without the fix the fp32 position is exactly 7.0, so floor = 7 and
901+
# the fraction is 0, misplacing the injection at cell 8 in 2D.
902+
u = TimeFunction(name='u', grid=grid, space_order=2, time_order=1)
903+
c = 0.6999999990000001
904+
sf.coordinates.data[0, :] = c
905+
sf.data[:] = 1.0
906+
907+
op = Operator(sf.inject(field=u.forward, expr=sf))
908+
op.apply(time_M=0)
909+
910+
# fp64-truth: nearly all mass at (7, 7), essentially none at (8, 8).
911+
assert np.isclose(u.data[1, 7, 7], 1.0, atol=1e-5)
912+
assert np.isclose(u.data[1, 8, 8], 0.0, atol=1e-10)
913+
873914

874915
# ---------------------------------------------------------------------------
875916
# Complex-valued interpolation

‎tests/test_mpi.py‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3350,11 +3350,15 @@ def test_adjoint_codegen(self, shape, kernel, space_order, save, mode):
33503350
solver = acoustic_setup(shape=shape, spacing=[15. for _ in shape], kernel=kernel,
33513351
tn=500, space_order=space_order, nrec=130,
33523352
preset='layers-isotropic', dtype=np.float64)
3353+
# Count only meaningful kernel/comm calls; ignore memory-management
3354+
# (posix_memalign/free) that grows with sparse-position preambles.
3355+
skip = {'posix_memalign', 'free'}
3356+
33533357
op_fwd = solver.op_fwd(save=save)
3354-
fwd_calls = FindNodes(Call).visit(op_fwd)
3358+
fwd_calls = [c for c in FindNodes(Call).visit(op_fwd) if c.name not in skip]
33553359

33563360
op_adj = solver.op_adj()
3357-
adj_calls = FindNodes(Call).visit(op_adj)
3361+
adj_calls = [c for c in FindNodes(Call).visit(op_adj) if c.name not in skip]
33583362

33593363
assert len(fwd_calls) == 1
33603364
assert len(adj_calls) == 1

0 commit comments

Comments
 (0)