Skip to content

Commit b1147aa

Browse files
committed
compiler: fix buffering corner case such as interp
1 parent e1ba6d6 commit b1147aa

7 files changed

Lines changed: 89 additions & 45 deletions

File tree

‎.github/workflows/lint.yaml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -88,6 +88,6 @@ jobs:
8888
for DOCKERFILE in docker/Dockerfile.*; \
8989
do \
9090
echo " Linting $DOCKERFILE"; \
91-
hadolint "$DOCKERFILE" \
91+
hadolint --ignore DL3066 "$DOCKERFILE" \
9292
|| exit 1; \
9393
done

‎devito/ir/equations/algorithms.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -365,6 +365,8 @@ def generate_conditionals(expr, input_expr, ordering):
365365
cond = d.relation(cond, GuardFactor(d))
366366
conditionals[d] = cond
367367

368+
if not conditionals and not input_expr.implicit_dims:
369+
return expr, conditionals
368370
# Merge conditionals when possible. E.g., if an implicit_dim shares
369371
# its parent Dimension with another ConditionalDimension, the two
370372
# conditions can be merged into a single guard.

‎devito/ir/stree/algorithms.py‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -169,7 +169,15 @@ def preprocess(clusters, options=None, **kwargs):
169169
for c in clusters:
170170
if c.is_halo_touch:
171171
hs = HaloScheme.union(e.rhs.halo_scheme for e in c.exprs)
172-
queue.append(c.rebuild(exprs=[], halo_scheme=hs))
172+
# Peel syncs (e.g. a WaitLock) off, they must survive an unbound halo.
173+
# Restrict the IterationSpace to the syncs' prefix so we don't carry
174+
# inner Dimensions (e.g. a layer's VirtualDimension) that have no
175+
# content left, which would yield empty guarded/iteration nodes
176+
if c.syncs:
177+
key = lambda d: any(d in s._defines for s in c.syncs) # noqa: B023
178+
ispace = c.ispace.prefix(key)
179+
processed.append(c.rebuild(exprs=[], ispace=ispace))
180+
queue.append(c.rebuild(exprs=[], syncs={}, halo_scheme=hs))
173181

174182
elif c.is_dist_reduce:
175183
processed.append(c)

‎devito/ir/support/guards.py‎

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -65,7 +65,14 @@ class GuardFactor(Guard, CondEq, Pickable):
6565

6666
__rargs__ = ('d',)
6767

68-
def __new__(cls, d, **kwargs):
68+
def __new__(cls, *args, **kwargs):
69+
if len(args) != 1:
70+
# Reconstruction with relational args (e.g. via sympy `_subs`): the
71+
# factor semantics no longer hold, so degrade to a plain relational
72+
base = CondNe if issubclass(cls, CondNe) else CondEq
73+
return base(*args, **kwargs)
74+
75+
d, = args
6976
assert d.is_Conditional
7077

7178
obj = super().__new__(cls, d.parent % d.symbolic_factor, 0)
@@ -169,6 +176,11 @@ def __new__(cls, d, index, direction,
169176
except TypeError:
170177
pass
171178

179+
return cls._new(p0, p1, d, index, direction, d_min=d_min, d_max=d_max, **kwargs)
180+
181+
@classmethod
182+
def _new(cls, p0, p1, d, index, direction, d_min=None, d_max=None):
183+
172184
obj = super().__new__(cls, p0, p1, evaluate=False)
173185

174186
obj.d = d
@@ -572,7 +584,9 @@ def pairwise_or(*guards):
572584

573585
@_uxreplace_handle.register(BaseGuardBoundNext)
574586
def _(expr, args, kwargs):
575-
return expr.func(expr.d, expr.index, expr.direction, **kwargs)
587+
p0, p1 = args
588+
return expr._new(p0, p1, expr.d, expr.index, expr.direction,
589+
**kwargs)
576590

577591

578592
@singledispatch

‎devito/passes/clusters/asynchrony.py‎

Lines changed: 22 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
WaitLock, WithLock, normalize_syncs
99
)
1010
from devito.passes.clusters.utils import in_critical_region, is_memcpy
11-
from devito.symbolics import IntDiv, retrieve_terminals, uxreplace
11+
from devito.symbolics import IntDiv, retrieve_dimensions, retrieve_terminals, uxreplace
1212
from devito.tools import OrderedSet, is_integer, timed_pass
1313
from devito.types import CustomDimension, Lock, VirtualDimension
1414

@@ -339,27 +339,29 @@ def _actions_from_update_memcpy(c, d, bounds, clusters, actions, sregistry):
339339

340340
syncs = {d: [
341341
ReleaseLock(handle, target),
342-
PrefetchUpdate(handle, target, tindex, function, findex, d,
342+
PrefetchUpdate(handle, target, tindex, f, findex, d,
343343
origin=e.rhs, gid=gid)
344344
]}
345345
syncs = {**c.syncs, **syncs}
346346

347347
pc = c.rebuild(exprs=expr, ispace=ispace, guards=guards, syncs=syncs)
348348

349-
# Since we're turning `e` into a prefetch, we need to:
350-
# 1) attach a WaitLock SyncOp to the first Cluster accessing `target`
351-
# 2) insert the prefetch Cluster right after the last Cluster accessing `target`
352-
# 3) drop the original Cluster performing a memcpy-like fetch
353-
n = clusters.index(c)
354-
first = None
355-
last = None
356-
for c1 in clusters[n+1:]:
357-
if target in c1.scope.reads:
358-
if first is None:
359-
first = c1
360-
last = c1
361-
assert first is not None
362-
assert last is not None
349+
# Wait before the first, prefetch after the last access to `target`, then
350+
# drop the memcpy `c`. `c` may be toposorted amid the readers, so scan them
351+
# all; count only reads over the streamed `pd`, not the buffer-init loop.
352+
reads = [c1 for c1 in clusters
353+
if c1 is not c and target in c1.scope.reads and pd in c1.ispace.itdims]
354+
assert reads
355+
first = reads[0]
356+
last = reads[-1]
357+
358+
# Advance `last` past its loop nest so the prefetch follows it rather than
359+
# splitting it, e.g. severing an interpolation's store from its point loop
360+
nest_dims = set(last.ispace.itdims) - set(d._defines)
361+
for c1 in clusters[clusters.index(last)+1:]:
362+
if not nest_dims & set(c1.ispace.itdims):
363+
break
364+
last = c1
363365

364366
actions[first].syncs[d].append(WaitLock(handle, target))
365367
actions[last].insert.append(pc)
@@ -433,7 +435,10 @@ def _(expr, dim, dir):
433435
434436
3//2 - 1 + 0 = 0
435437
"""
436-
if expr.lhs._defines & dim._defines:
438+
dims = retrieve_dimensions(expr.lhs)
439+
assert len(dims) == 1
440+
ldim = dims.pop()
441+
if ldim._defines & dim._defines:
437442
if dir == 1:
438443
return expr + dir
439444
else:

‎devito/passes/clusters/buffering.py‎

Lines changed: 38 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -311,7 +311,11 @@ def _optimize(self, clusters, descriptors):
311311
processed.append(c)
312312
continue
313313

314-
key1 = lambda d: not d._defines & v.dim._defines # noqa: B023
314+
# Lift only along the buffer's own spatial Dimensions, not
315+
# unrelated ones (e.g. sparse points of an indirect read)
316+
key1 = lambda d: (not d._defines & v.dim._defines and # noqa: B023
317+
any(d._defines & bd._defines # noqa: B023
318+
for bd in v.bdims)) # noqa: B023
315319
dims = c.ispace.project(key1).itdims
316320
ispace = c.ispace.lift(dims, key0())
317321
processed.append(c.rebuild(ispace=ispace))
@@ -331,8 +335,8 @@ def _make_task_groups(self, descriptors):
331335
return stamps
332336

333337
task_sets = (
334-
[b for b, v in descriptors.items() if v.is_writeonly],
335-
[b for b, v in descriptors.items() if v.is_readonly],
338+
[b for b, v in descriptors.items() if any(vb.is_writeonly for vb in v)],
339+
[b for b, v in descriptors.items() if any(vb.is_readonly for vb in v)],
336340
)
337341

338342
for tasks in task_sets:
@@ -519,30 +523,37 @@ def itdims(self):
519523
def ispace(self):
520524
# The IterationSpace within which the buffer will be accessed
521525

522-
# NOTE: The `key` is to avoid Clusters including `f` but not directly
523-
# using it in an expression, such as HaloTouch Clusters
524-
def key(c):
525-
bufferdim = any(i in c.ispace.dimensions for i in self.bdims)
526-
xd_only = all(d._defines & self.xd._defines for d in c.ispace.dimensions)
527-
return bufferdim or xd_only
528-
529526
ispaces = set()
527+
indirect = set()
530528
for c in self.clusters:
531-
if not key(c):
532-
continue
533-
534-
# Skip wild clusters (e.g. HaloTouch Clusters)
529+
# Skip wild clusters (e.g. HaloTouch Clusters), which include `f`
530+
# but do not directly use it in an expression
535531
if c.is_wild:
536532
continue
537533

538-
# Iterations space and buffering dims
539-
edims = [d for d in self.bdims if d not in c.ispace.dimensions]
540-
if not edims:
541-
ispaces.add(c.ispace)
542-
else:
543-
# Add all missing buffering dimensions and reorder to
544-
# avoid duplicates with different ordering
545-
ispaces.add(c.ispace.insert(self.dim, edims).reorder())
534+
bufferdim = any(i in c.ispace.dimensions for i in self.bdims)
535+
xd_only = all(d._defines & self.xd._defines for d in c.ispace.dimensions)
536+
537+
if bufferdim or xd_only:
538+
# `c` iterates (at least some of) the buffer's spatial dims
539+
edims = [d for d in self.bdims if d not in c.ispace.dimensions]
540+
if not edims:
541+
ispaces.add(c.ispace)
542+
else:
543+
# Add all missing buffering dimensions and reorder to
544+
# avoid duplicates with different ordering
545+
ispaces.add(c.ispace.insert(self.dim, edims).reorder())
546+
elif ((self.f in c.scope.reads or self.f in c.scope.writes) and
547+
self.dim.root in c.ispace.dimensions):
548+
# `c` accesses `f` indirectly (e.g. interpolation), so it doesn't
549+
# iterate `bdims`; span the buffer's own Dimensions instead
550+
tispace = c.ispace.project(lambda i: i._defines & self.dim.root._defines)
551+
indirect.add(tispace.insert(self.dim, list(self.bdims)))
552+
553+
# Indirect accessors define the ispace only for a read-only streamed
554+
# buffer, where nothing iterates the buffer's own Dimensions directly
555+
if not ispaces:
556+
ispaces = indirect
546557

547558
if len(ispaces) > 1:
548559
# Best effort to make buffering work in the presence of multiple
@@ -663,7 +674,11 @@ def last_idx(self):
663674
mapper = {}
664675
func = vmax if self.is_forward_buffering else vmin
665676
for c in self.lastwrite + self.firstread:
666-
indices = extract_indices(self.f, self.dim, [c])
677+
# Consider all Clusters sharing `c`'s guards, so the leading edge is
678+
# found even when `f` is accessed at different offsets across them
679+
# (e.g. `f` and `f.forward` in separate Eqs)
680+
group = [c1 for c1 in self.clusters if c1.guards == c.guards]
681+
indices = extract_indices(self.f, self.dim, group)
667682
idx = func(*[Vector(i) for i in indices])[0]
668683
mapper[c] = idx
669684

‎tests/test_gpu_common.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1419,7 +1419,7 @@ def test_streaming_complete(self):
14191419
assert len([i for i in FindSymbols().visit(op2) if i.is_Array]) == 9 - diff
14201420
# The streaming/buffering cannot be fused since the guards of
14211421
# vas and vasb/vasf are not compatible.
1422-
assert len(op3._func_table) == 14 - diff
1422+
assert len(op3._func_table) == 12 - diff
14231423
assert len([i for i in FindSymbols().visit(op3) if i.is_Array]) == 9 - diff
14241424

14251425
op0.apply(time_m=15, time_M=35, save_shift=0)

0 commit comments

Comments
 (0)