99from devito .finite_differences import EvalDerivative , IndexDerivative , Weights
1010from devito .ir import (
1111 PARALLEL_IF_PVT , SEPARABLE , SEQUENTIAL , Cluster , ClusterGroup , ExprGeometry , Forward ,
12- Interval , IntervalGroup , IterationSpace , LabeledVector , Queue , Vector , extrema ,
13- maximum , minimum , normalize_properties , relax_properties , unbounded , vmax , vmin
12+ Interval , IntervalGroup , IterationSpace , LabeledVector , Properties , Queue , Vector ,
13+ extrema , maximum , minimum , normalize_properties , relax_properties , unbounded , vmax ,
14+ vmin
1415)
1516from devito .passes .clusters .cse import _cse
1617from devito .passes .clusters .utils import expose_tuning_knobs
1920 uxreplace
2021)
2122from devito .tools import (
22- Reconstructable , Stamp , as_mapper , as_tuple , flatten , frozendict , generator ,
23- is_integer , split , timed_pass
23+ Reconstructable , Stamp , as_mapper , as_tuple , flatten , generator , is_integer , split ,
24+ timed_pass
2425)
2526from devito .types import (
2627 CustomDimension , Eq , Hyperplane , IncrDimension , Indexed , ModuloDimension , Size ,
@@ -288,6 +289,7 @@ def _do_generate(self, exprs, exclude, cbk_search, cbk_compose=None):
288289 free_symbols = i .free_symbols
289290 if {a .function for a in free_symbols } & exclude :
290291 continue
292+
291293 mapper .add (i , make , terms )
292294
293295 return mapper
@@ -304,6 +306,31 @@ def __init__(self, sregistry, options, platform):
304306 def process (self , clusters ):
305307 return self ._process_fatd (clusters , 1 , xtracted = [])
306308
309+ @classmethod
310+ def _make_exclude (cls , clusters , d , p ):
311+ """
312+ The symbols an extraction must not touch.
313+ """
314+ # Rule out extractions that would break data dependencies
315+ exclude = set ().union (* [c .scope .writes for c in clusters ])
316+
317+ # Rule out extractions that depend on the Dimension currently investigated,
318+ # as they clearly wouldn't be invariants
319+ exclude .update ({d , * p .sub_iterators })
320+
321+ # An extraction is hoisted out of `d`, but it inherits its guard as it
322+ # is, so it must not be hoisted past any Dimension the guard reads, or
323+ # the guard would be evaluated where it is not defined. Excluding those
324+ # Dimensions keeps the extraction inside the loops defining them. A
325+ # SEQUENTIAL Dimension is exempt, its loop enclosing the extraction's own
326+ for c in clusters :
327+ exclude .update (
328+ i for i in c .guards .dimensions
329+ if not c .properties .is_sequential (i ._defines )
330+ )
331+
332+ return exclude
333+
307334 def callback (self , clusters , prefix , xtracted = None ):
308335 if not prefix :
309336 return clusters
@@ -314,12 +341,7 @@ def callback(self, clusters, prefix, xtracted=None):
314341 if d .is_Virtual :
315342 return clusters
316343
317- # Rule out extractions that would break data dependencies
318- exclude = set ().union (* [c .scope .writes for c in clusters ])
319-
320- # Rule out extractions that depend on the Dimension currently investigated,
321- # as they clearly wouldn't be invariants
322- exclude .update ({d , * p .sub_iterators })
344+ exclude = self ._make_exclude (clusters , d , p )
323345
324346 key = lambda c : self ._lookup_key (c , d )
325347 processed = list (clusters )
@@ -343,7 +365,7 @@ def callback(self, clusters, prefix, xtracted=None):
343365 def _lookup_key (self , c , d ):
344366 ispace = c .ispace .reset ()
345367 intervals = c .ispace .intervals .drop (d ).reset ()
346- properties = frozendict ({d : relax_properties (v ) for d , v in c .properties .items ()})
368+ properties = Properties ({d : relax_properties (v ) for d , v in c .properties .items ()})
347369
348370 return AliasKey (ispace , intervals , c .dtype , c .guards , properties )
349371
0 commit comments