Skip to content

Commit 3a0ad97

Browse files
committed
compiler: Revamp SubDimension DDA
1 parent 3c0f3ea commit 3a0ad97

8 files changed

Lines changed: 612 additions & 117 deletions

File tree

‎devito/ir/support/basic.py‎

Lines changed: 121 additions & 90 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from sympy import Expr, S
88

99
from devito.ir.support.space import Backward, null_ispace
10-
from devito.ir.support.utils import AccessMode, extrema
10+
from devito.ir.support.utils import AccessMode, erange, extrema
1111
from devito.ir.support.vector import LabeledVector, Vector
1212
from devito.symbolics import (
1313
compare_ops, q_affine, q_comp_acc, q_constant, retrieve_indexed, search
@@ -358,6 +358,9 @@ def distance(self, other, logical=False):
358358
# E.g., `uv(x).x` and `uv(x).y` -- not a real dependence!
359359
return Vector(S.ImaginaryUnit)
360360

361+
if disjoint_subdims(self, other):
362+
return Vector(S.ImaginaryUnit)
363+
361364
ret = []
362365
for sit, oit in zip(self.itintervals, other.itintervals, strict=False):
363366
n = len(ret)
@@ -369,20 +372,14 @@ def distance(self, other, logical=False):
369372
# E.g., `self=R<f,[x]>` and `self.itintervals=(x, i)`
370373
break
371374

372-
# If over SubDimensions, check disjointness
373-
test = disjoint_subdims(self[n], other[n], sai, oai, sit, oit)
374-
if test == DISJOINT:
375-
return Vector(S.ImaginaryUnit)
376-
elif test == MAYBE_OVERLAP:
377-
ret.append(S.Infinity)
378-
continue
379-
380375
try:
381376
if not (sit == oit and sai.root is oai.root):
382377
# E.g., `self=R<f,[x + 2]>` and `other=W<f,[i + 1]>`
383378
# E.g., `self=R<f,[x]>`, `other=W<f,[x + 1]>`,
384379
# `self.itintervals=(x<0>,)`, `other.itintervals=(x<1>,)`
385-
return vinf(ret)
380+
# Keep looking: a later axis may prove disjointness
381+
ret.append(S.Infinity)
382+
continue
386383
except AttributeError:
387384
# E.g., `self=R<f,[cy]>` and `self.itintervals=(y,)` => `sai=None`
388385
pass
@@ -1152,7 +1149,29 @@ def reads_smart_gen(self, f):
11521149
"""
11531150
Generate all read accesses to a given function.
11541151
1155-
StencilDimensions, if any, are replaced with their extrema.
1152+
StencilDimensions, if any, are replaced with:
1153+
1154+
* in presence of SubDimensions: the range of points they span;
1155+
* in all other cases: just their extrema, since it suffices to
1156+
capture all possible dependencies.
1157+
1158+
The reason SubDimensions must be treated specially -- with a full set
1159+
of TimedAccess objects getting generated -- is to handle the special
1160+
case of slabs thinner than the stencil’s reach. For example, consider
1161+
the following scenario:
1162+
1163+
* A SubDimension with just two points, 10 and 11;
1164+
* One equation writes `F[10]` and `F[11]`;
1165+
* Another equation runs over the same SubDimension reading the stencil
1166+
`F[x-4] ... F[x+4]`.
1167+
1168+
If we examine only the two extreme stencil offsets:
1169+
1170+
* `F[x-4]` reads points 6–7: no overlap.
1171+
* `F[x+4]` reads points 14–15: no overlap.
1172+
1173+
But interior offsets certainly overlap -- for instance, `F[x-1]` reads
1174+
9–10, which includes the producer’s point 10.
11561175
11571176
Notes
11581177
-----
@@ -1163,9 +1182,16 @@ def reads_smart_gen(self, f):
11631182
be found. For example, a DiscreteFunction would never appear among
11641183
the iteration symbols.
11651184
"""
1185+
uses_subdims = lambda i: any(d.is_Sub for d in i.ispace.dimensions)
1186+
11661187
if isinstance(f, (Function, Temp, TempArray, TBArray)):
11671188
for i in self.getreads(f):
1168-
for j in extrema(i.access):
1189+
if uses_subdims(i):
1190+
expand = erange
1191+
else:
1192+
expand = extrema
1193+
1194+
for j in expand(i.access):
11691195
yield TimedAccess(j, i.mode, i.timestamp, i.ispace)
11701196

11711197
else:
@@ -1581,90 +1607,95 @@ def skippable_interval(d, ispace, it):
15811607
return d is None or (d in ispace and not d._defines & it.dim._defines)
15821608

15831609

1584-
# Possible return values for `disjoint_subdims`
1585-
INAPPLICABLE = 0
1586-
DISJOINT = 1
1587-
MAYBE_OVERLAP = 2
1588-
1589-
1590-
def disjoint_subdims(e0, e1, d0, d1, it0, it1):
1610+
def disjoint_subdims(a0, a1):
15911611
"""
1592-
Determine whether two accesses span distinct pieces of the same
1593-
SubDimension decomposition.
1594-
1595-
Consider a root Dimension `x` with bounds `x_m` and `x_M`. A valid
1596-
left/middle/right decomposition with thicknesses `L` and `R` is::
1597-
1598-
xl = [x_m, x_m + L - 1]
1599-
xm = [x_m + L, x_M - R]
1600-
xr = [x_M - R + 1, x_M]
1601-
1602-
These intervals are pairwise disjoint. Replacing `xl`, `xm`, or `xr`
1603-
with `x` in an affine access removes the choice of partition piece while
1604-
retaining the relative access. If two such normalized accesses have zero
1605-
distance, they apply the same affine map to disjoint intervals and therefore
1606-
cannot refer to the same data point. The apparent dependence is imaginary.
1607-
1608-
For example, `f[xl]` and `f[xm]` normalize to `f[x]` and `f[x]`;
1609-
they are independent. The same holds for `f[xl + 1]` and `f[xm + 1]`
1610-
when their iteration intervals have equal offsets. By contrast, `f[xl]`
1611-
and `f[xm - 1]` normalize to different accesses, and the latter may reach
1612-
into the left piece, so they must be treated conservatively.
1613-
1614-
This proof requires distinct pieces of the same root, compatible declared
1615-
thicknesses, affine accesses, and iteration intervals with equal offsets and
1616-
directions. Runtime bounds are assumed to preserve the declared partition.
1617-
Return DISJOINT if disjointness is proven, and MAYBE_OVERLAP if the
1618-
intervals are aligned SubDimensions but are not proven disjoint. In
1619-
particular, two declarations of the same left, right, or middle piece
1620-
overlap along this Dimension. MAYBE_OVERLAP lets the caller record an
1621-
infinite distance and inspect later Dimensions, which may still prove the
1622-
multidimensional accesses disjoint. Return INAPPLICABLE if this test does not
1623-
apply, so that the general distance analysis can classify the dependence.
1612+
Determine whether two TimedAccesses touch disjoint SubDimension regions
1613+
of the same Function.
1614+
1615+
Compare symbolic accessed bounds, including shifts and stencil points.
1616+
Block intervals are promoted to their logical SubDimensions. Bounds and
1617+
thicknesses remain symbolic: MPI decomposition and runtime overrides can
1618+
change their values independently of the defaults.
1619+
1620+
For example, `xl = [m, m + L - 1]` and `xm = [m + L, M - R]` are
1621+
disjoint when they share the symbol `L`, whatever its runtime value.
1622+
Equal default thicknesses alone do not establish that relationship.
1623+
1624+
Opposite left/right slabs are assumed to form a valid partition: their
1625+
thicknesses satisfy `L + R <= N`. For translated stencil accesses, the
1626+
interior must also accommodate their combined inward reach. For example,
1627+
a pointwise left write and a right read at offset -4 require four interior
1628+
points. Runtime space_order checks cover explicit middle SubDimensions,
1629+
not arbitrary left/right pairs; no concrete domain size or thickness is
1630+
used here.
1631+
1632+
Match data axes independently of the iteration nests. Return True if any
1633+
axis proves separation, False otherwise. Accesses over the same interval
1634+
use the general distance analysis.
16241635
"""
1625-
try:
1626-
# E.g., `f[xl]` over `(xl,)` and `f[xm]` over `(xm,)` need this
1627-
# special test, while accesses over the same `(xl,)` should use general
1628-
# distance analysis, so we can return immediately in such a case
1629-
if not (d0.is_Sub and
1630-
d1.is_Sub and
1631-
d0.root is d1.root and
1632-
it0.dim.root is d0.root and
1633-
it1.dim.root is d1.root and
1636+
for e0, e1, d0, d1 in zip(a0, a1, a0.aindices, a1.aindices, strict=False):
1637+
it0 = a0.intervals[d0]
1638+
it1 = a1.intervals[d1]
1639+
if it0.is_Null or it1.is_Null:
1640+
continue
1641+
1642+
it0 = it0.promote(lambda d: d.is_Incr)
1643+
it1 = it1.promote(lambda d: d.is_Incr)
1644+
if not (it0.dim.is_Sub and
1645+
it1.dim.is_Sub and
1646+
it0.dim.root is it1.dim.root and
16341647
it0 != it1):
1635-
return INAPPLICABLE
1636-
except AttributeError:
1637-
return INAPPLICABLE
1638-
1639-
if (d0.is_left and d1.is_middle) or \
1640-
(d0.is_middle and d1.is_left):
1641-
is_partition = d0.ltkn.value == d1.ltkn.value
1642-
elif (d0.is_middle and d1.is_right) or \
1643-
(d0.is_right and d1.is_middle):
1644-
is_partition = d0.rtkn.value == d1.rtkn.value
1645-
elif d0.is_left and d1.is_right:
1646-
is_partition = d0.ltkn.value is not None and d1.rtkn.value is not None
1647-
elif d0.is_right and d1.is_left:
1648-
is_partition = d0.rtkn.value is not None and d1.ltkn.value is not None
1649-
else:
1650-
is_partition = False
1651-
1652-
if not is_partition:
1653-
return MAYBE_OVERLAP
1654-
1655-
if not q_affine(e0, d0) or not q_affine(e1, d1):
1656-
return MAYBE_OVERLAP
1648+
continue
16571649

1658-
if it0.offsets != it1.offsets or it0.direction is not it1.direction:
1659-
return MAYBE_OVERLAP
1650+
bounds = []
1651+
for e, d, it in ((e0, d0, it0), (e1, d1, it1)):
1652+
if not q_affine(e, d):
1653+
break
16601654

1661-
e0 = e0._subs(d0, d0.root)
1662-
e1 = e1._subs(d1, d1.root)
1655+
lower, upper = [], []
1656+
for v in erange(e):
1657+
slope = v.diff(d)
1658+
if slope.is_nonnegative:
1659+
m, M = it.symbolic_min, it.symbolic_max
1660+
elif slope.is_nonpositive:
1661+
M, m = it.symbolic_min, it.symbolic_max
1662+
else:
1663+
break
1664+
lower.append(v._subs(d, m))
1665+
upper.append(v._subs(d, M))
1666+
else:
1667+
bounds.append((sympy.Min(*lower), sympy.Max(*upper)))
1668+
1669+
if len(bounds) == 2:
1670+
(m0, M0), (m1, M1) = bounds
1671+
mapper = {}
1672+
1673+
dl, dr = (it0.dim, it1.dim) if it0.dim.is_left else (it1.dim, it0.dim)
1674+
dlp, drp = dl.parent, dr.parent
1675+
1676+
if dl.is_left and dr.is_right and dlp is drp:
1677+
# A valid partition satisfies L + R <= N, where N is the parent
1678+
# extent; an explicit middle SubDomain checks this at construction.
1679+
# Further, for stencils, we require that:
1680+
# `N - L - R >= the combined inward reach`
1681+
# so accesses from opposite slabs cannot meet. Explicit middle
1682+
# SubDimensions check for at least space_order interior points
1683+
# at *op.apply time*, accounting for runtime overrides. Without
1684+
# an explicit middle, the gap assumption is unchecked
1685+
gap = sympy.Dummy(nonnegative=True)
1686+
if e0.diff(d0) == e1.diff(d1) == 1:
1687+
M, m = (M0, m1) if it0.dim.is_left else (M1, m0)
1688+
reach = (M - dl.symbolic_max - m + dr.symbolic_min).expand()
1689+
if is_integer(reach):
1690+
gap += max(0, reach)
1691+
1692+
mapper[dlp.symbolic_max] = dlp.symbolic_min + dl.ltkn + dr.rtkn + gap - 1
1693+
1694+
if (M0 - m1).subs(mapper).is_negative or \
1695+
(M1 - m0).subs(mapper).is_negative:
1696+
return True
16631697

1664-
if e0 - e1 == 0:
1665-
return DISJOINT
1666-
else:
1667-
return MAYBE_OVERLAP
1698+
return False
16681699

16691700

16701701
def disjoint_test(e0, e1, d, it):

‎devito/operator/operator.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -716,7 +716,7 @@ def _prepare_arguments(self, autotune=None, estimate_memory=False, **kwargs):
716716
except AttributeError:
717717
pass
718718
if d.is_Derived:
719-
d._arg_check(args)
719+
d._arg_check(args, **kwargs)
720720

721721
# Turn arguments into a format suitable for the generated code
722722
# E.g., instead of NumPy arrays for Functions, the generated code expects

‎devito/passes/iet/orchestration.py‎

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -125,9 +125,9 @@ def _make_syncarray(self, iet, sync_ops, layer):
125125
def _make_prefetchupdate(self, iet, sync_ops, layer, wrap=True):
126126
return self._make_async_task(prefetchupdate, iet, sync_ops, layer, wrap)
127127

128-
@iet_pass
129-
def process(self, iet):
130-
callbacks = {
128+
@property
129+
def _callbacks(self):
130+
return {
131131
WaitLock: self._make_waitlock,
132132
WithLock: self._make_withlock,
133133
SyncArray: self._make_syncarray,
@@ -139,13 +139,15 @@ def process(self, iet):
139139
AsyncCallable: self._make_async_callable
140140
}
141141

142+
@iet_pass
143+
def process(self, iet):
142144
# Collect the compatible asynchronous task groups, if any
143145
task_groups = TaskGroups()
144146
if self.npthreads:
145147
CollectTasks(task_groups).visit(iet)
146148

147149
# Lower the SyncSpots in a single bottom-up traversal, atomically lowering
148-
lowerer = LowerSyncSpots(callbacks, task_groups, self.sregistry)
150+
lowerer = LowerSyncSpots(self._callbacks, task_groups, self.sregistry)
149151
iet = lowerer.visit(iet)
150152

151153
return iet, {'efuncs': lowerer.efuncs}

‎devito/types/dimension.py‎

Lines changed: 41 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -818,6 +818,46 @@ def _arg_values(self, interval, grid=None, **kwargs):
818818
# themselves
819819
return {}
820820

821+
def _arg_check(self, args, *_args, **kwargs):
822+
# These modules depend on Dimension, so importing them above would cycle
823+
from devito.mpi import mpi_raise # noqa: PLC0415
824+
from devito.symbolics import subs_op_args # noqa: PLC0415
825+
826+
if not self.is_middle:
827+
return
828+
829+
# Function._arg_check visits original axes (e.g. x_ltkn), whereas `args`
830+
# contains the concretized thicknesses (x_ltkn0, ...). The Operator checks
831+
# the matching concrete SubDimensions separately in its dimension loop
832+
if self not in args.op.dimensions:
833+
return
834+
835+
d = self.root
836+
if args.grid is not None and args.grid.is_distributed(d):
837+
# Check the global runtime region: its MPI-local slices may be empty
838+
# or smaller than space_order even for a non-degenerate global interior
839+
size = args.grid.size_map[d].glb
840+
values = {**args,
841+
d.min_name: kwargs.get(d.min_name, 0),
842+
d.max_name: kwargs.get(d.max_name, kwargs.get(d.name, size - 1)),
843+
**{t.name: kwargs.get(t.name, t.value) for t in self.thickness}}
844+
else:
845+
values = args
846+
size = int(subs_op_args(self.symbolic_size, values))
847+
848+
# Runtime overrides do not change the compiled stencil order
849+
items = [f.space_order for f in args.op.input if f.is_DiscreteFunction]
850+
space_order = max(items, default=0)
851+
852+
if size < space_order:
853+
error = (f"Expected at least {space_order} interior points along "
854+
f"`{self.parent}` (space_order), but runtime arguments leave {size}")
855+
else:
856+
error = None
857+
858+
comm = args.comm if args.options['mpi'] else None
859+
mpi_raise(error, InvalidArgument, comm=comm)
860+
821861

822862
class MultiSubDimension(AbstractSubDimension):
823863

@@ -1436,7 +1476,7 @@ def _arg_values(self, interval, grid=None, args=None, **kwargs):
14361476
# Avoid OOB (will end up here only in case of tiny iteration spaces)
14371477
return {name: 1}
14381478

1439-
def _arg_check(self, args, *_args):
1479+
def _arg_check(self, args, *_args, **kwargs):
14401480
try:
14411481
name = self.step.name
14421482
except AttributeError:

‎devito/types/grid.py‎

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -641,12 +641,6 @@ def __subdomain_finalize_legacy__(self, grid):
641641
try:
642642
# Case ('middle', int, int)
643643
side, ltkn, rtkn = v
644-
if side != 'middle':
645-
raise ValueError(f"Expected side 'middle', not `{side}`")
646-
sub_dimensions.append(SubDimension.middle(f'i{k.name}',
647-
k, ltkn, rtkn))
648-
thickness = s-ltkn-rtkn
649-
sdshape.append(thickness)
650644
except ValueError:
651645
side, thickness = v
652646
constructor = {'left': SubDimension.left,
@@ -663,6 +657,23 @@ def __subdomain_finalize_legacy__(self, grid):
663657
) from None
664658
sub_dimensions.append(constructor(f'i{k.name}', k, thickness))
665659
sdshape.append(thickness)
660+
else:
661+
if side != 'middle':
662+
raise ValueError(f"Expected side 'middle', not `{side}`")
663+
664+
# A `middle` region expects `ltkn + rtkn <= s` in the global Grid.
665+
# This ensures that the left and right regions won't overlap
666+
thickness = s-ltkn-rtkn
667+
if thickness < 0:
668+
raise ValueError(
669+
f"SubDomain `{self.name}` has combined thickness "
670+
f"{ltkn + rtkn} along `{k}`, exceeding the Grid size {s}"
671+
)
672+
673+
sub_dimensions.append(
674+
SubDimension.middle(f'i{k.name}', k, ltkn, rtkn)
675+
)
676+
sdshape.append(thickness)
666677

667678
self._shape = tuple(sdshape)
668679
self._dimensions = tuple(sub_dimensions)

0 commit comments

Comments
 (0)