Skip to content

Commit 327c61f

Browse files
committed
compiler: Simplify HALO-write handling
1 parent 6adba1a commit 327c61f

7 files changed

Lines changed: 182 additions & 164 deletions

File tree

‎devito/ir/clusters/algorithms.py‎

Lines changed: 15 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
from devito.ir.support import (
1515
Any, Backward, Forward, IterationSpace, Scope, detect_halo_writes, erange, pull_dims
1616
)
17+
from devito.logger import warning
1718
from devito.mpi.halo_scheme import HaloScheme, HaloTouch
1819
from devito.mpi.reduction_scheme import DistReduce
1920
from devito.symbolics import limits_mapper, retrieve_indexed, uxreplace, xreplace_indices
@@ -632,21 +633,22 @@ def _update(reductions):
632633

633634
def check_halo_writes(clusters):
634635
"""
635-
Reject explicit HALO writes along Dimensions split across MPI ranks.
636+
Warn about HALO writes along Dimensions not fixed to 1 in the Grid topology.
636637
"""
637638
for c in clusters:
638-
dims = set()
639-
for f in c.scope.writes:
640-
if not f.is_DiscreteFunction or f.grid is None:
641-
continue
642-
dist = f.grid.distributor
643-
dims.update(d.root for d, n in zip(dist.dimensions, dist.topology,
644-
strict=True) if n > 1)
645-
646-
key = lambda d: d.root in dims # noqa: B023
647-
if dims and detect_halo_writes(c, key):
648-
raise CompilationError("Cannot write to the HALO along distributed "
649-
"Dimensions")
639+
try:
640+
grid = c.grid
641+
except ValueError:
642+
grid = None
643+
644+
topology = {}
645+
if grid is not None and grid.topology is not None:
646+
topology = dict(zip(grid.dimensions, grid.topology, strict=True))
647+
648+
key = lambda d: d in c.dist_dimensions and topology.get(d) != 1 # noqa: B023
649+
if detect_halo_writes(c, key):
650+
warning("Writing to the HALO along potentially distributed Dimensions; "
651+
"set their Grid topology entries to 1")
650652

651653

652654
def normalize(clusters, sregistry=None, options=None, platform=None, **kwargs):

‎devito/ir/support/basic.py‎

Lines changed: 21 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -505,67 +505,38 @@ def touched_nodomain(self, findex):
505505
if not self.affine(findex):
506506
return (False, False)
507507

508-
index = self[findex]
509508
d = self.aindices[findex]
510-
511-
if d is None:
512-
index_min = index_max = index
513-
else:
514-
try:
515-
m, M = self.intervals[d].offsets
516-
except KeyError:
517-
return (False, False)
518-
519-
coefficient = index.diff(d)
520-
if coefficient.is_positive:
521-
index_min = index.subs(d, d.symbolic_min + m)
522-
index_max = index.subs(d, d.symbolic_max + M)
523-
elif coefficient.is_negative:
524-
index_min = index.subs(d, d.symbolic_max + M)
525-
index_max = index.subs(d, d.symbolic_min + m)
526-
elif coefficient.is_zero:
527-
index_min = index_max = index
528-
else:
509+
limits = []
510+
if d is not None:
511+
i = self.intervals[d]
512+
if i.is_Null:
529513
return (False, False)
514+
limits.append((d, d.symbolic_min + i.lower, d.symbolic_max + i.upper))
530515

531-
size_nodomain_left = self.function._size_nodomain[findex].left
532-
domain_min = findex.symbolic_min + size_nodomain_left
533-
domain_max = findex.symbolic_max + size_nodomain_left
516+
# Runtime DOMAIN bounds may select any part of the allocated extent
517+
for v in (findex.symbolic_min, findex.symbolic_max):
518+
if v.is_Symbol:
519+
limits.append((v, S.Zero, findex.symbolic_size - 1))
534520

535-
def bound(expr, maximize):
521+
def outside(expr):
522+
# A negative maximum distance proves the entire access is outside
536523
expr = sympy.expand(expr)
537-
538-
size = findex.symbolic_size
539-
limits = [
540-
(findex.symbolic_min, S.Zero, size - 1),
541-
(findex.symbolic_max, S.Zero, size - 1)
542-
]
543-
544524
for symbol, lower, upper in limits:
545-
if not symbol.is_Symbol:
546-
continue
547-
548525
coefficient = expr.diff(symbol)
549526
if coefficient.has(symbol):
550-
return None
551-
elif coefficient.is_positive:
552-
value = upper if maximize else lower
553-
elif coefficient.is_negative:
554-
value = lower if maximize else upper
555-
elif coefficient.is_zero:
556-
continue
527+
return False
528+
elif coefficient.is_nonnegative:
529+
expr = expr.subs(symbol, upper)
530+
elif coefficient.is_nonpositive:
531+
expr = expr.subs(symbol, lower)
557532
else:
558-
return None
559-
560-
expr = expr.subs(symbol, value)
561-
562-
return expr
533+
return False
563534

564-
left = bound(index_max - domain_min, True)
565-
right = bound(index_min - domain_max, False)
535+
return expr.is_negative is True
566536

567-
return (left is not None and left.is_negative is True,
568-
right is not None and right.is_positive is True)
537+
index = self[findex] - self.function._size_nodomain[findex].left
538+
return (outside(index - findex.symbolic_min),
539+
outside(findex.symbolic_max - index))
569540

570541
def touched_halo(self, findex):
571542
"""

‎devito/passes/clusters/buffering.py‎

Lines changed: 29 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -9,8 +9,7 @@
99
from devito.exceptions import CompilationError
1010
from devito.ir import (
1111
Backward, Cluster, Forward, GuardBound, GuardFactor, InitArray, Interval,
12-
IntervalGroup, IterationSpace, Properties, Queue, Vector, detect_halo_writes,
13-
lower_exprs, vmax, vmin
12+
IterationSpace, Properties, Queue, Vector, detect_halo_writes, lower_exprs, vmax, vmin
1413
)
1514
from devito.logger import warning
1615
from devito.passes.clusters.utils import is_memcpy
@@ -496,61 +495,61 @@ def expand_halo_transfers(clusters, mapper):
496495
writes. For example, `usave` in `Eq(usave, u)` must eventually receive `u`'s
497496
populated HALO if a preceding `Eq` writes into `u`'s HALO.
498497
"""
499-
buffered = {f for f, _ in mapper}
500-
if not buffered:
498+
if not mapper:
501499
return clusters
502500

501+
# Get HALO writes along the buffered dimensions
503502
bdims = set()
504503
for b in mapper.values():
505504
bdims.update(d for d in b.dimensions if not isinstance(d, BufferDimension))
506-
key = lambda d: d in bdims
507505

508506
halo_writes = set()
509507
for c in clusters:
510-
for w in detect_halo_writes(c, key):
508+
for w in detect_halo_writes(c, bdims.__contains__):
511509
halo_writes.add(w.function)
510+
if not halo_writes:
511+
return clusters
512512

513+
# Expand the IterationSpace over the necessary amount of HALO; in doing so,
514+
# check the expanded footprint of every access, including shifted reads.
515+
# Writes must be pointwise so that the whole destination halo is filled
516+
buffered = {f for f, _ in mapper}
513517
processed = []
514518
for c in clusters:
515-
scope = c.scope
516-
targets = set(scope.writes) & buffered
517-
if c.is_wild or not targets or not halo_writes.intersection(scope.reads):
519+
writes = c.scope.writes_tensor
520+
521+
if c.is_wild or \
522+
writes.isdisjoint(buffered) or \
523+
halo_writes.isdisjoint(c.scope.reads):
518524
processed.append(c)
519525
continue
520526

521-
if scope.writes_tensor != targets:
527+
if not writes <= buffered:
522528
raise CompilationError(
523529
"Cannot expand a mixed Cluster over the halo while buffering"
524530
)
525531

526532
ispace = c.ispace
527-
for f in targets:
533+
for f in writes:
528534
ispace = _include_halo(ispace, f)
529535

530-
# Check the expanded footprint of every access, including shifted reads.
531-
# Writes must be pointwise so that the whole destination halo is filled
532-
for a in scope.accesses:
536+
for a in c.scope.accesses:
533537
f = a.function
534-
if not f.is_AbstractFunction:
535-
continue
536-
537-
for d in f.dimensions:
538-
if not key(d):
539-
continue
540-
if d not in ispace.dimensions:
541-
raise CompilationError(
542-
f"Cannot expand access to `{f.name}` over the halo"
543-
)
544538

539+
for d in bdims.intersection(a.findices):
545540
size = f._size_nodomain[d]
546-
offset = simplify(a[d] - d)
547-
if not is_integer(offset) or (a.is_write and offset != size.left):
541+
offset = simplify(a[d] - d - size.left)
542+
543+
if d not in ispace.dimensions or \
544+
not is_integer(offset) or \
545+
(a.is_write and offset != 0):
548546
raise CompilationError(
549-
f"Cannot expand non-pointwise access to `{f.name}` over the halo"
547+
f"Cannot expand access to `{f.name}` over the halo"
550548
)
551549

552550
i = ispace[d]
553-
if i.lower + offset < 0 or i.upper + offset > sum(size):
551+
if i.lower + offset < -size.left or \
552+
i.upper + offset > size.right:
554553
raise CompilationError(
555554
f"Insufficient halo for `{f.name}` in buffered write"
556555
)
@@ -562,10 +561,10 @@ def expand_halo_transfers(clusters, mapper):
562561

563562
def _include_halo(ispace, f):
564563
"""Extend `ispace` to include `f`'s HALO."""
565-
ihalo = IntervalGroup([
564+
ihalo = [
566565
Interval(i.dim, -f._size_halo[i.dim].left, f._size_halo[i.dim].right, i.stamp)
567566
for i in ispace if i.dim in f.dimensions
568-
])
567+
]
569568

570569
return IterationSpace.union(ispace, IterationSpace(ihalo))
571570

‎devito/types/array.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -520,7 +520,7 @@ def initvalue(self):
520520
'_mem_rvalue', '__padding_dtype__', '_size_domain', '_size_halo',
521521
'_size_owned', '_size_padding', '_size_nopad', '_size_nodomain',
522522
'_offset_domain', '_offset_halo', '_offset_owned',
523-
'_dist_dimensions', '_C_get_field', 'grid',
523+
'_dist_dimensions', '_decomposition', '_C_get_field', 'grid',
524524
*AbstractFunction.__properties__):
525525
locals()[i] = property(lambda self, v=i: getattr(self.c0, v))
526526

‎tests/test_buffering.py‎

Lines changed: 19 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -94,18 +94,18 @@ def test_write_only_with_halo_source(forward):
9494
hx, hy = usave._size_halo.left[1:]
9595
for t in range(nt-1):
9696
assert np.all(usave.data[t] == t + forward)
97-
for i in range(1, 5):
98-
actual = usave.data_with_halo[t, hx:hx + grid.shape[0], hy - i]
99-
assert np.all(actual == -(t + forward))
97+
actual = usave.data_with_halo[t, hx:hx + grid.shape[0], hy-4:hy]
98+
assert np.all(actual == -(t + forward))
10099

101100

102101
@pytest.mark.parametrize('space_order, shift', [(0, 0), (8, -1), (8, 1), (10, 1)])
102+
@switchconfig(autopadding=False)
103103
def test_write_only_with_halo_source_bounds(space_order, shift):
104104
grid = Grid(shape=(17, 17))
105105
y = grid.dimensions[-1]
106106

107107
u = TimeFunction(name='u', grid=grid, space_order=8)
108-
v = TimeFunction(name='v', grid=grid, space_order=space_order, padding=0)
108+
v = TimeFunction(name='v', grid=grid, space_order=space_order)
109109
usave = TimeFunction(name='usave', grid=grid, space_order=8, save=5)
110110

111111
k = CustomDimension(name='k', parent=y, symbolic_min=1,
@@ -116,7 +116,7 @@ def test_write_only_with_halo_source_bounds(space_order, shift):
116116
Eq(usave, u.forward + v.forward._subs(y, y + shift))]
117117

118118
if space_order == 10:
119-
# An extra halo point accommodates the shifted read
119+
# A wider halo accommodates the shifted read
120120
v.data_with_halo[:] = 2
121121
op = Operator(eqns, opt='buffering', name='save_shifted_halo')
122122
op.apply(time_M=3)
@@ -128,7 +128,8 @@ def test_write_only_with_halo_source_bounds(space_order, shift):
128128
Operator(eqns, opt='buffering')
129129

130130

131-
def test_halo_transfers_non_time_dimension():
131+
@pytest.mark.parametrize('mixed', [False, True])
132+
def test_halo_transfers_non_time_dimension(mixed):
132133
s = Dimension(name='s')
133134
x = Dimension(name='x')
134135
u = Function(name='u', dimensions=(s, x), shape=(5, 17),
@@ -142,9 +143,18 @@ def test_halo_transfers_non_time_dimension():
142143

143144
mirror = Cluster(lower_exprs(Eq(u[s+1, -k], -u[s+1, k])),
144145
IterationSpace([Interval(s), Interval(k)]))
145-
save = Cluster(lower_exprs(Eq(usave[s, x], u[s+1, x])),
146-
IterationSpace([Interval(s), Interval(x)]))
147-
clusters = expand_halo_transfers([mirror, save], {(usave, save.guards): b})
146+
eqns = [Eq(usave[s, x], u[s+1, x])]
147+
if mixed:
148+
eqns.append(Eq(u[s, x], 0))
149+
save = Cluster(lower_exprs(eqns), IterationSpace([Interval(s), Interval(x)]))
150+
mapper = {(usave, save.guards): b}
151+
152+
if mixed:
153+
with pytest.raises(CompilationError, match='mixed Cluster'):
154+
expand_halo_transfers([mirror, save], mapper)
155+
return
156+
157+
clusters = expand_halo_transfers([mirror, save], mapper)
148158

149159
assert clusters[0] is mirror
150160
assert clusters[1].ispace[x].offsets == (-4, 4)

0 commit comments

Comments
 (0)