Skip to content

Commit b6933db

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 b6933db

10 files changed

Lines changed: 455 additions & 100 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: 51 additions & 26 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):
@@ -2960,10 +2971,13 @@ def test_fullopt(self):
29602971
assert summary0[('section1', None)].ops == 44
29612972
assert np.isclose(summary0[('section0', None)].oi, 3.136, atol=0.001)
29622973

2963-
assert summary1[('section0', None)].ops == 31
2964-
assert summary1[('section1', None)].ops == 88
2965-
assert summary1[('section2', None)].ops == 25
2966-
assert np.isclose(summary1[('section0', None)].oi, 1.767, atol=0.001)
2974+
# section0 is now the sparse-position preamble hoisted out of the
2975+
# time loop by `CireInvariantsSparse`; the dense stencil is section1.
2976+
assert summary1[('section0', None)].ops == 9
2977+
assert summary1[('section1', None)].ops == 31
2978+
assert summary1[('section2', None)].ops == 40
2979+
assert summary1[('section3', None)].ops == 25
2980+
assert np.isclose(summary1[('section1', None)].oi, 0.884, atol=0.001)
29672981

29682982
assert np.allclose(u0.data, u1.data, atol=10e-5)
29692983
assert np.allclose(rec0.data, rec1.data, atol=10e-5)
@@ -3022,9 +3036,13 @@ def test_fullopt(self):
30223036
assert np.allclose(self.tti_noopt[0].data, v.data, atol=10e-1)
30233037
assert np.allclose(self.tti_noopt[1].data, rec.data, atol=10e-1)
30243038

3025-
# Check expected opcount/oi
3026-
assert summary[('section1', None)].ops == 92
3027-
assert np.isclose(summary[('section1', None)].oi, 1.99, atol=0.001)
3039+
# Check expected opcount/oi. The dense TTI stencil moved from
3040+
# `section1` to `section2` (inside multipass0) because sparse-position
3041+
# preambles now occupy the leading sections. Access the sops directly
3042+
# via the profiler's subsection map.
3043+
op_fwd = wavesolver.op_fwd()
3044+
stencil = op_fwd._profiler._subsections['multipass0']['section2']
3045+
assert stencil.sops == 92
30283046

30293047
# With optimizations enabled, there should be exactly four BlockDimensions
30303048
op = wavesolver.op_fwd()
@@ -3035,15 +3053,15 @@ def test_fullopt(self):
30353053
assert y.parent is y0_blk0
30363054
assert not x._defines & y._defines
30373055

3038-
# Also, in this operator, we expect six temporary Arrays:
3039-
# * all of the six Arrays are allocated on the heap
3040-
# * with OpenMP:
3041-
# four Arrays are defined globally for the cos/sin temporaries
3042-
# 3 Arrays are defined globally for the sparse positions temporaries
3043-
# and two additional bock-sized Arrays are defined locally
3056+
# Temporary Arrays expected:
3057+
# * 4 global for the cos/sin temporaries
3058+
# * 6 global for the sparse-position preambles: per grid dim, one
3059+
# fp64 `pos` and one int32 `int(floor(pos))`, coalesced across
3060+
# the src+rec inject/interp.
3061+
# * 2 block-local
30443062
arrays = [i for i in FindSymbols().visit(op) if i.is_Array]
30453063
extra_arrays = 2
3046-
assert len(arrays) == 4 + extra_arrays
3064+
assert len(arrays) == 4 + 6 + extra_arrays
30473065
assert all(i._mem_heap and not i._mem_external for i in arrays)
30483066
bns, pbs = assert_blocking(op, {'x0_blk0'})
30493067

@@ -3076,19 +3094,26 @@ def test_fullopt_w_mpi(self, mode):
30763094
])
30773095
def test_opcounts(self, space_order, expected):
30783096
op = self.tti_operator(opt='advanced', space_order=space_order)
3079-
sections = list(op.op_fwd()._profiler._sections.values())
3080-
assert sections[1].sops == expected
3097+
# The dense TTI stencil moved inside `multipass0` when sparse-position
3098+
# preambles took the leading section slots.
3099+
op_fwd = op.op_fwd()
3100+
stencil = op_fwd._profiler._subsections['multipass0']['section2']
3101+
assert stencil.sops == expected
30813102

30823103
@switchconfig(profiling='advanced')
30833104
@pytest.mark.parametrize('space_order,exp_ops,exp_arrays', [
3084-
(4, 122, 6), (8, 225, 7)
3105+
(4, 122, 12), (8, 225, 13)
30853106
])
30863107
def test_opcounts_adjoint(self, space_order, exp_ops, exp_arrays):
30873108
wavesolver = self.tti_operator(space_order=space_order,
30883109
opt=('advanced', {'openmp': False}))
30893110
op = wavesolver.op_adj()
30903111

3091-
assert op._profiler._sections['section1'].sops == exp_ops
3112+
# Dense stencil is nested inside `multipass0` (sparse-position
3113+
# preambles occupy the leading sections). `exp_arrays` grows by 6 to
3114+
# include the sparse `pos`/`int_floor` preambles (2 per grid dim).
3115+
stencil = op._profiler._subsections['multipass0']['section2']
3116+
assert stencil.sops == exp_ops
30923117
assert len([i for i in FindSymbols().visit(op) if i.is_Array]) == exp_arrays
30933118

30943119

‎tests/test_gradient.py‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -57,7 +57,13 @@ def test_gradient_checkpointing(self, dtype, opt, space_order, kernel, shape, sp
5757
gradient2, _ = wave.jacobian_adjoint(residual, u0, vp=v0, checkpointing=False,
5858
grad=grad)
5959

60-
assert np.allclose(gradient.data, gradient2.data, atol=0, rtol=0)
60+
# Not strictly bit-identical since sparse-position preambles now
61+
# compute `(c - o)/h` in fp64: the fp64 intermediates may be reordered
62+
# or fused differently across the two Operator instantiations, giving
63+
# a sub-ulp fp32 drift near the source. Tolerate a small mismatch.
64+
peak = float(np.max(np.abs(gradient2.data)))
65+
atol = 1e-5 * peak if dtype == np.float32 else 1e-12 * peak
66+
assert np.allclose(gradient.data, gradient2.data, atol=atol, rtol=0)
6167

6268
@skipif(['cpu64-icc', 'cpu64-arm'])
6369
@pytest.mark.parametrize('tn', [750.])

0 commit comments

Comments
 (0)