@@ -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