Skip to content

Commit 3fcd485

Browse files
committed
compiler: Expose separate attribute in SubDomain API
1 parent 4ec3dc6 commit 3fcd485

3 files changed

Lines changed: 73 additions & 2 deletions

File tree

‎devito/types/grid.py‎

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -583,6 +583,14 @@ class SubDomain(AbstractSubDomain):
583583
region of ``d_size - (N + M)`` points starting at ``N`` and finishing
584584
at ``d_sizeM - M``.
585585
586+
Attributes
587+
----------
588+
separated : bool, default=True
589+
Require a stencil-safe gap between opposite left/right regions, as in
590+
SubDimension. Set to False to allow touching or overlapping regions.
591+
Applies to SubDimensions generated from tuple definitions; explicit
592+
Dimensions returned by :meth:`define` retain their own settings.
593+
586594
Examples
587595
--------
588596
An "Inner" SubDomain, which spans the entire domain except for an exterior
@@ -614,6 +622,8 @@ class SubDomain(AbstractSubDomain):
614622
especially when defining BCs.
615623
"""
616624

625+
separated = True
626+
617627
def __subdomain_finalize__(self):
618628
self.__subdomain_finalize_legacy__(self.grid)
619629
self._distributor = SubDistributor(self)
@@ -655,7 +665,8 @@ def __subdomain_finalize_legacy__(self, grid):
655665
f"Maximum thickness of dimension {k.name} "
656666
f"is {s}, not {thickness}"
657667
) from None
658-
sub_dimensions.append(constructor(f'i{k.name}', k, thickness))
668+
sub_dimensions.append(constructor(f'i{k.name}', k, thickness,
669+
separated=self.separated))
659670
sdshape.append(thickness)
660671
else:
661672
if side != 'middle':
@@ -671,7 +682,8 @@ def __subdomain_finalize_legacy__(self, grid):
671682
)
672683

673684
sub_dimensions.append(
674-
SubDimension.middle(f'i{k.name}', k, ltkn, rtkn)
685+
SubDimension.middle(f'i{k.name}', k, ltkn, rtkn,
686+
separated=self.separated)
675687
)
676688
sdshape.append(thickness)
677689

‎tests/test_pickle.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,10 @@ def define(self, dimensions):
4949
return {x: x, y: ('middle', 1, 1), z: ('right', 2)}
5050

5151

52+
class OverlappingSD(SD):
53+
separated = False
54+
55+
5256
@pytest.mark.parametrize('pickle', [pickle0, pickle1])
5357
class TestBasic:
5458

@@ -112,6 +116,16 @@ def test_enrichedtuple_rebuild(self, pickle):
112116
assert new_t.left == tup.left
113117
assert new_t.right == tup.right
114118

119+
def test_subdomain(self, pickle):
120+
grid = Grid(shape=(5, 5, 5))
121+
sd = OverlappingSD(grid=grid)
122+
123+
new_sd = pickle.loads(pickle.dumps(sd))
124+
125+
assert not new_sd.separated
126+
assert new_sd.shape == sd.shape
127+
assert all(not d.separated for d in new_sd.dimensions if d.is_Sub)
128+
115129
@pytest.mark.parametrize('on_sd', [False, True])
116130
def test_function(self, pickle, on_sd):
117131
grid = Grid(shape=(3, 3, 3))

‎tests/test_subdomains.py‎

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -145,6 +145,31 @@ def make_region():
145145
expected = size - 16 if spec[0] == 'middle' else 16
146146
assert make_region().shape == (expected,)
147147

148+
@pytest.mark.parametrize('legacy', [False, True])
149+
@pytest.mark.parametrize('separated', [False, True])
150+
def test_separated(self, legacy, separated):
151+
class Region(SubDomain):
152+
name = 'region'
153+
154+
def define(self, dimensions):
155+
x, y, z = dimensions
156+
return {x: ('left', 4), y: ('middle', 1, 1), z: ('right', 4)}
157+
158+
class OverlappingRegion(Region):
159+
separated = False
160+
161+
cls = Region if separated else OverlappingRegion
162+
if legacy:
163+
region = cls()
164+
Grid(shape=(7, 7, 7), subdomains=(region,))
165+
else:
166+
region = cls(grid=Grid(shape=(7, 7, 7)))
167+
168+
assert region.separated is separated
169+
assert all(d.separated is separated for d in region.dimensions)
170+
assert all(t.separated is separated for d in region.dimensions
171+
for t in d.thickness)
172+
148173
def test_definitions(self):
149174

150175
class sd0(SubDomain):
@@ -289,6 +314,26 @@ def define(self, dimensions):
289314

290315
assert_structure(op, ['t', 'txyz', 'txyz'], 'txyzyz')
291316

317+
def test_overlapping(self):
318+
class Overlapping(ReducedDomain):
319+
separated = False
320+
321+
grid = Grid(shape=(7, 5))
322+
left = Overlapping(('left', 4), ('middle', 1, 1), grid=grid)
323+
right = Overlapping(('right', 4), None, grid=grid)
324+
325+
f = Function(name='f', grid=grid)
326+
327+
eqs = [Eq(f, f + 1, subdomain=left), Eq(f, 2*f + 2, subdomain=right)]
328+
329+
op = Operator(eqs, name='overlapping_subdomains')
330+
op.apply()
331+
332+
expected = np.zeros(grid.shape)
333+
expected[:4, 1:-1] += 1
334+
expected[3:, :] = 2*expected[3:, :] + 2
335+
assert np.array_equal(f.data, expected)
336+
292337

293338
class TestSubDomainScheduling:
294339
"""Tests scheduling across different SubDomains."""

0 commit comments

Comments
 (0)