1+ import contextlib
12from itertools import groupby
23
34from sympy import sympify
@@ -67,7 +68,11 @@ def blocking(clusters, sregistry, options):
6768 clusters = AnalyzeSkewing ().process (clusters )
6869
6970 if options ['blocklevels' ] > 0 :
70- clusters = SynthesizeBlocking (sregistry , options ).process (clusters )
71+ if options ['blockrelax' ] == 'device-aware' :
72+ synthesizer = SynthesizeBlockingDeviceAware (sregistry , options )
73+ else :
74+ synthesizer = SynthesizeBlocking (sregistry , options )
75+ clusters = synthesizer .process (clusters )
7176
7277 if options ['skewing' ]:
7378 clusters = SynthesizeSkewing (options ).process (clusters )
@@ -162,6 +167,14 @@ def __init__(self, options):
162167
163168 self .gpu_fit = options .get ('gpu-fit' , ())
164169
170+ def _process_fatd (self , clusters , level , prefix = None ):
171+ processed = []
172+ for _ , group in groupby (clusters , key = lambda c : c .ispace ):
173+ g = list (group )
174+ processed .extend (Queue ._process_fatd (self , g , level , prefix ))
175+
176+ return processed
177+
165178 def _make_key_hook (self , cluster , level ):
166179 return (is_on_device (cluster .functions , self .gpu_fit ),)
167180
@@ -315,18 +328,21 @@ def callback(self, clusters, prefix):
315328 return processed
316329
317330
318- class SynthesizeBlocking (Queue ):
331+ class SynthesizeBlockingBase (Queue ):
332+
333+ mapper = None
334+ """
335+ A mapping from a tuple of (Dimension, number of stencil points) to a tuple
336+ of BlockDimensions, so that we can reuse existing BlockDimensions to avoid
337+ unnecessary `steps`. Disabled by default, to be enabled in subclasses that
338+ need it (e.g., SynthesizeBlocking).
339+ """
319340
320341 def __init__ (self , sregistry , options ):
321342 self .sregistry = sregistry
322343
323344 self .levels = options ['blocklevels' ]
324345
325- # Track the BlockDimensions created so far so that we can reuse them
326- # in case of Clusters that are different but share the same number of
327- # stencil points
328- self .mapper = {}
329-
330346 super ().__init__ ()
331347
332348 def process (self , clusters ):
@@ -340,11 +356,11 @@ def _make_key_guards(self, cluster, ispace):
340356 if not cluster .properties .is_blockable (i .dim ))
341357
342358 def _derive_block_dims (self , clusters , prefix , d ):
343- # Can I reuse existing BlockDimensions to avoid a proliferation of steps?
359+ # Can I reuse existing BlockDimensions to avoid unnecessary ` steps` ?
344360 k = stencil_footprint (clusters , d )
345361 try :
346362 return self .mapper [k ]
347- except KeyError :
363+ except ( KeyError , TypeError ) :
348364 pass
349365
350366 base = self .sregistry .make_name (prefix = d .root .name )
@@ -362,7 +378,10 @@ def _derive_block_dims(self, clusters, prefix, d):
362378 bd = BlockDimension (d .name , bd , bd , bd + bd .step - 1 , 1 , size = step )
363379 block_dims .append (bd )
364380
365- retval = self .mapper [k ] = tuple (block_dims ), bd
381+ retval = tuple (block_dims ), bd
382+
383+ with contextlib .suppress (TypeError ):
384+ self .mapper [k ] = retval
366385
367386 return retval
368387
@@ -404,6 +423,28 @@ def callback(self, clusters, prefix):
404423 return processed
405424
406425
426+ class SynthesizeBlocking (SynthesizeBlockingBase ):
427+
428+ def __init__ (self , sregistry , options ):
429+ # Track the BlockDimensions created so far so that we can reuse them
430+ # in case of Clusters that are different but share the same number of
431+ # stencil points
432+ self .mapper = {}
433+
434+ super ().__init__ (sregistry , options )
435+
436+
437+ class SynthesizeBlockingDeviceAware (SynthesizeBlockingBase ):
438+
439+ def _process_fdta (self , clusters , level , prefix = None ):
440+ processed = []
441+ for _ , group in groupby (clusters , key = lambda c : c .ispace ):
442+ g = list (group )
443+ processed .extend (Queue ._process_fdta (self , g , level , prefix ))
444+
445+ return processed
446+
447+
407448def stencil_footprint (clusters , d ):
408449 """
409450 Compute the number of stencil points in the given Dimension `d` across the
0 commit comments