Skip to content

Commit 49855a0

Browse files
committed
compiler: Tweak blocking for CPML-like subdomains and GPUs
1 parent 25c2dc1 commit 49855a0

2 files changed

Lines changed: 53 additions & 11 deletions

File tree

‎devito/passes/clusters/blocking.py‎

Lines changed: 51 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import contextlib
12
from itertools import groupby
23

34
from 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+
407448
def stencil_footprint(clusters, d):
408449
"""
409450
Compute the number of stencil points in the given Dimension `d` across the

‎tests/test_dle.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -264,7 +264,8 @@ def test_leftright_subdims(self):
264264

265265
op = Operator(eqns, opt=('fission', 'blocking', {'blockrelax': 'device-aware'}))
266266

267-
bns, _ = assert_blocking(op, {'x0_blk0', 'x1_blk0', 'x2_blk0'})
267+
bns, _ = assert_blocking(op, {'x0_blk0', 'x1_blk0', 'x2_blk0',
268+
'x3_blk0', 'x4_blk0'})
268269
assert all(IsPerfectIteration().visit(i) for i in bns.values())
269270
assert all(len(FindNodes(Iteration).visit(i)) == 4 for i in bns.values())
270271

0 commit comments

Comments
 (0)