Skip to content

Commit d88604f

Browse files
committed
tests: Force blocking on MPI test to ensure it also works on CPU
1 parent 30ee739 commit d88604f

1 file changed

Lines changed: 12 additions & 13 deletions

File tree

‎tests/test_mpi.py‎

Lines changed: 12 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -1149,7 +1149,7 @@ def define(self, dimensions):
11491149

11501150
grid = Grid(shape=(21, 21, 21), extent=(1., 1., 1.))
11511151
domains = {spec: SpeccedDomain(f"x{spec[0]}_y{spec[1]}_z{spec[2]}", nb, spec,
1152-
grid=grid)
1152+
grid=grid)
11531153
for spec in specs}
11541154

11551155
p = TimeFunction(name='p', grid=grid, space_order=so, time_order=to)
@@ -1159,34 +1159,33 @@ def define(self, dimensions):
11591159
# Fields live on the base grid, not `v` -- `subdomain=v` is passed
11601160
# explicitly on each Eq instead.
11611161
psi = VectorTimeFunction(name=f"psi_{v.name}", grid=grid, space_order=so,
1162-
time_order=to, staggered=(None, None, None))
1162+
time_order=to, staggered=(None, None, None))
11631163
zeta = VectorTimeFunction(name=f"zeta_{v.name}", grid=grid, space_order=so,
1164-
time_order=to, staggered=(None, None, None))
1164+
time_order=to, staggered=(None, None, None))
11651165
psi_eqs.append(Eq(psi, 1, subdomain=v))
11661166

11671167
# "Diagonal" derivative pattern -- component `i` of zeta is the
11681168
# derivative of component `i` of psi along dimension `i`. This
11691169
# specific pattern is required to trigger the bug; grad()/div()/a
11701170
# plain vector add alone do not.
11711171
zeta_diag = VectorTimeFunction([getattr(psi[i], f"d{d.name}")
1172-
for i, d in enumerate(grid.dimensions)])
1172+
for i, d in enumerate(grid.dimensions)])
11731173
zeta_eqs.append(Eq(zeta, zeta_diag, subdomain=v))
11741174
p_eqs.append(Eq(p.forward, psi.div(), subdomain=v))
11751175

1176-
op = Operator(psi_eqs + zeta_eqs + p_eqs)
1176+
op = Operator(psi_eqs + zeta_eqs + p_eqs,
1177+
opt=('advanced', {'blockrelax': 'device-aware'}))
11771178
op.apply(time_M=1)
11781179

11791180
# No misplaced halo exchanges: every halo-exchange node must sit above
11801181
# (never inside) any blocking Iteration.
11811182
incr_iterations = [i for i in FindNodes(Iteration).visit(op) if i.dim.is_Incr]
1182-
assert incr_iterations, "no blocking Iterations found -- check would be vacuous"
1183-
1184-
if configuration['mpi']:
1185-
# HaloSpots have already been lowered into concrete Calls by mpiize()
1186-
halo_types = (HaloUpdateCall, HaloUpdateList)
1187-
else:
1188-
# HaloSpots remain in the IET as transparent (no-op) wrappers
1189-
halo_types = (HaloSpot,)
1183+
assert incr_iterations, "no blocking Iterations found"
1184+
1185+
# HaloSpots have already been lowered into concrete Calls by mpiize() in the MPI
1186+
# case, but HaloSpots remain in the IET as transparent (no-op) wrappers in the
1187+
# serial case.
1188+
halo_types = (HaloUpdateCall, HaloUpdateList) if configuration['mpi'] else (HaloSpot,)
11901189

11911190
for i in incr_iterations:
11921191
assert len(FindNodes(halo_types).visit(i)) == 0

0 commit comments

Comments
 (0)