Skip to content

Commit f335bb1

Browse files
Merge pull request #2973 from devitocodes/simplify-fission
compiler: Simplify fission
2 parents c9b9cfa + 72a5fdf commit f335bb1

5 files changed

Lines changed: 72 additions & 37 deletions

File tree

‎devito/ir/clusters/cluster.py‎

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -248,6 +248,14 @@ def is_phase_marker(self):
248248
def is_critical_region(self):
249249
return self._is_type(CriticalRegion)
250250

251+
@cached_property
252+
def is_thread_rendezvous(self):
253+
"""
254+
True if it contains a synchronization point at which all participating
255+
threads must arrive before any may proceed.
256+
"""
257+
return self.is_thread_pool_sync and not self.is_thread_wait
258+
251259
@cached_property
252260
def is_thread_pool_sync(self):
253261
return self._is_type(ThreadPoolSync)
@@ -655,6 +663,12 @@ def __hash__(self):
655663
def concatenate(cls, *cgroups):
656664
return list(chain(*cgroups))
657665

666+
def rebuild(self, **kwargs):
667+
clusters = kwargs.get('clusters', self)
668+
ispace = kwargs.get('ispace', self.ispace)
669+
670+
return self.__class__(clusters, ispace=ispace)
671+
658672
@cached_property
659673
def exprs(self):
660674
return flatten(c.exprs for c in self)
@@ -663,6 +677,10 @@ def exprs(self):
663677
def scope(self):
664678
return Scope(exprs=self.exprs)
665679

680+
@cached_property
681+
def functions(self):
682+
return self.scope.functions
683+
666684
@cached_property
667685
def ispace(self):
668686
return self._ispace

‎devito/passes/clusters/aliases.py‎

Lines changed: 1 addition & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
maximum, minimum, normalize_properties, relax_properties, unbounded, vmax, vmin
1414
)
1515
from devito.passes.clusters.cse import _cse
16+
from devito.passes.clusters.utils import expose_tuning_knobs
1617
from devito.symbolics import (
1718
Uxmapper, estimate_cost, retrieve_functions, reuse_if_untouched, search, sympy_dtype,
1819
uxreplace
@@ -1080,27 +1081,6 @@ def optimize_clusters_msds(clusters):
10801081
return processed
10811082

10821083

1083-
def expose_tuning_knobs(clusters, sregistry):
1084-
"""
1085-
Replace all pre-existing BlockDimensions with fresh ones, to enable
1086-
separate tuning for the CIRE-generated temporaries.
1087-
"""
1088-
# Create the new BlockDimensions
1089-
callback = lambda i: sregistry.make_name(prefix=i)
1090-
1091-
mapper = {}
1092-
for d in set().union(*[c.used_dimensions for c in clusters]):
1093-
if d.is_Block:
1094-
mapper.update(d._rebuild_hierarchy(callback))
1095-
1096-
if not mapper:
1097-
return clusters
1098-
1099-
processed = [c.subs(mapper) for c in clusters]
1100-
1101-
return processed
1102-
1103-
11041084
def pick_best(variants):
11051085
"""
11061086
Return the variant with the best theoretical performance.

‎devito/passes/clusters/misc.py‎

Lines changed: 11 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,10 @@
11
from itertools import groupby, product
22

33
from devito.ir.clusters import Queue, cluster_pass
4-
from devito.ir.support import SEPARABLE, SEQUENTIAL, Scope
4+
from devito.ir.support import SEPARABLE, Scope
55
from devito.passes.clusters.utils import in_critical_region
66
from devito.symbolics import pow_to_mul
7-
from devito.tools import Stamp, flatten, frozendict, timed_pass
7+
from devito.tools import Stamp, flatten, timed_pass
88
from devito.types import Hyperplane
99

1010
__all__ = ['Lift', 'fission', 'optimize_hyperplanes', 'optimize_pows']
@@ -123,7 +123,7 @@ def callback(self, clusters, prefix):
123123
d = prefix[-1].dim
124124

125125
# Do not waste time if definitely illegal
126-
if any(SEQUENTIAL in c.properties[d] for c in clusters):
126+
if any(c.properties.is_sequential(d) for c in clusters):
127127
return clusters
128128

129129
# Do not waste time if definitely nothing to do
@@ -132,21 +132,21 @@ def callback(self, clusters, prefix):
132132

133133
# Analyze and abort if fissioning would break a dependence
134134
scope = Scope(flatten(c.exprs for c in clusters))
135-
if any(d._defines & dep.cause or dep.is_reduce(d) for dep in scope.d_all_gen()):
135+
if any(d._defines & dep.cause or dep.is_reduce(d) or dep.is_local
136+
for dep in scope.d_all_gen()):
136137
return clusters
137138

138139
processed = []
139-
for (it, guards), g in groupby(clusters, key=lambda c: self._key(c, prefix)):
140+
for it, g in groupby(clusters, key=lambda c: self._key(c, prefix)):
140141
group = list(g)
141142

142143
try:
143-
test0 = any(SEQUENTIAL in c.properties[it.dim] for c in group)
144+
test0 = any(c.properties.is_sequential(it.dim) for c in group)
144145
except AttributeError:
145-
# `it` is None because `c`'s IterationSpace has no `d` Dimension,
146-
# hence `key = (it, guards) = (None, guards)`
146+
# `it` is None because `c`'s IterationSpace has no `d` Dimension
147147
test0 = True
148148

149-
if test0 or guards:
149+
if test0:
150150
# Heuristic: no gain from fissioning if unable to ultimately
151151
# increase the number of collapsible iteration spaces, hence give up
152152
processed.extend(group)
@@ -161,14 +161,10 @@ def callback(self, clusters, prefix):
161161
def _key(self, c, prefix):
162162
try:
163163
index = len(prefix)
164-
dims = tuple(i.dim for i in prefix)
165-
166164
it = c.ispace[index]
167-
guards = frozendict({d: v for d, v in c.guards.items() if d in dims})
168-
169-
return (it, guards)
165+
return it
170166
except IndexError:
171-
return (None, c.guards)
167+
return None
172168

173169

174170
@timed_pass()

‎devito/passes/clusters/utils.py‎

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,8 @@
22
from devito.tools import as_tuple
33
from devito.types import CriticalRegion, Eq, Symbol
44

5-
__all__ = ['in_critical_region', 'is_memcpy', 'make_critical_sequence']
5+
__all__ = ['expose_tuning_knobs', 'in_critical_region', 'is_memcpy',
6+
'make_critical_sequence']
67

78

89
def is_memcpy(expr):
@@ -50,3 +51,24 @@ def in_critical_region(cluster, clusters):
5051
elif c.is_critical_region:
5152
maybe_found = c
5253
return None
54+
55+
56+
def expose_tuning_knobs(clusters, sregistry):
57+
"""
58+
Replace all pre-existing BlockDimensions with fresh ones, to enable
59+
separate tuning for the CIRE-generated temporaries.
60+
"""
61+
# Create the new BlockDimensions
62+
callback = lambda i: sregistry.make_name(prefix=i)
63+
64+
mapper = {}
65+
for d in set().union(*[c.used_dimensions for c in clusters]):
66+
if d.is_Block:
67+
mapper.update(d._rebuild_hierarchy(callback))
68+
69+
if not mapper:
70+
return clusters
71+
72+
processed = [c.subs(mapper) for c in clusters]
73+
74+
return processed

‎tests/test_fission.py‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -80,6 +80,25 @@ def test_nofission_as_illegal():
8080
assert_structure(op, ['t,x,y', 't,x,y'], 't,x,y,y')
8181

8282

83+
def test_nofission_local_scalar_dependence():
84+
"""
85+
Test there's no fission across a local scalar dependence.
86+
"""
87+
grid = Grid(shape=(3, 3))
88+
time = grid.time_dim
89+
x, y = grid.dimensions
90+
91+
g = Function(name='g', grid=grid, dtype=np.int32, space_order=0)
92+
h = Function(name='h', grid=grid, space_order=0)
93+
94+
eqns = [Eq(y.symbolic_max, g[x, 0], implicit_dims=(time, x)),
95+
Eq(h, y, implicit_dims=(time, x, y))]
96+
97+
op = Operator(eqns, opt='fission')
98+
99+
assert_structure(op, ['t,x', 't,x,y'], 't,x,y')
100+
101+
83102
def test_fission_partial():
84103
"""
85104
Test there's no fission if no increase in number of collapsible loops.

0 commit comments

Comments
 (0)