@@ -212,7 +212,8 @@ def test_read_only_w_offset():
212212 assert np .all (v .data == v1 .data )
213213
214214
215- def test_read_only_backwards ():
215+ @pytest .mark .parametrize ('async_degree,expected_size' , [(None , 3 ), (4 , 4 )])
216+ def test_read_only_backwards (async_degree , expected_size ):
216217 nt = 10
217218 grid = Grid (shape = (2 , 2 ))
218219
@@ -226,12 +227,14 @@ def test_read_only_backwards():
226227 eqns = [Eq (v .backward , v + u .backward + u + u .forward + 1. )]
227228
228229 op0 = Operator (eqns , opt = 'noop' )
229- op1 = Operator (eqns , opt = 'buffering' )
230+ op1 = Operator (eqns , opt = ('buffering' ,
231+ {'buf-async-degree' : async_degree }))
230232
231233 # Check generated code
232234 assert len (retrieve_iteration_tree (op1 )) == 4
233235 buffers = [i for i in FindSymbols ().visit (op1 ) if i .is_Array and i ._mem_heap ]
234236 assert len (buffers ) == 1
237+ assert buffers .pop ().symbolic_shape [0 ] == expected_size
235238
236239 op0 .apply (time_m = 1 )
237240 op1 .apply (time_m = 1 , v = v1 )
@@ -270,32 +273,83 @@ def test_read_only_backwards_unstructured():
270273 assert np .all (v .data == v1 .data )
271274
272275
273- @pytest .mark .parametrize ('async_degree' , [2 , 4 ])
274- def test_async_degree (async_degree ):
276+ @pytest .mark .parametrize ('async_degree' , [1 , 2 , 4 ])
277+ @pytest .mark .parametrize ('backward' , [False , True ],
278+ ids = ['forward' , 'backward' ])
279+ def test_async_degree (async_degree , backward ):
275280 nt = 10
276281 grid = Grid (shape = (4 , 4 ))
277282
278283 u = TimeFunction (name = 'u' , grid = grid , save = nt )
279284 u1 = TimeFunction (name = 'u' , grid = grid , save = nt )
280285
281- eqn = Eq (u .forward , u + 1 )
286+ lhs = u .backward if backward else u .forward
287+ eqn = Eq (lhs , u + 1 )
282288
283289 op0 = Operator (eqn , opt = 'noop' )
284290 op1 = Operator (eqn , opt = ('buffering' , {'buf-async-degree' : async_degree }))
285291
286292 # Check generated code
287293 assert len (retrieve_iteration_tree (op1 )) == 3
288- buffers = [i for i in FindSymbols ().visit (op1 ) if i .is_Array and i ._mem_heap ]
294+ buffers = [i for i in FindSymbols ().visit (op1 )
295+ if i .is_Array and i ._mem_heap ]
289296 assert len (buffers ) == 1
290- assert buffers .pop ().symbolic_shape [0 ] == async_degree
297+ assert buffers .pop ().symbolic_shape [0 ] == max ( 2 , async_degree )
291298
292- op0 .apply (time_M = nt - 2 )
293- op1 .apply (time_M = nt - 2 , u = u1 )
299+ kwargs = {'time_m' : 1 } if backward else {'time_M' : nt - 2 }
300+ op0 .apply (** kwargs )
301+ op1 .apply (u = u1 , ** kwargs )
294302
295303 assert np .all (u .data == u1 .data )
296304
297305
298- def test_two_homogeneous_buffers ():
306+ @pytest .mark .parametrize ('backward,expected_bounds' , [
307+ pytest .param (False , (0 , 8 ), id = 'forward' ),
308+ pytest .param (True , (1 , 9 ), id = 'backward' )
309+ ])
310+ @pytest .mark .parametrize ('async_degree' , [0 , 1 , 4 , 16 ])
311+ def test_async_degree_read_only (backward , expected_bounds , async_degree ):
312+ nt = 10
313+ grid = Grid (shape = (4 , 4 ))
314+
315+ u = TimeFunction (name = 'u' , grid = grid , save = nt )
316+ v = TimeFunction (name = 'v' , grid = grid )
317+ v1 = TimeFunction (name = 'v' , grid = grid )
318+
319+ u .data [:] = np .arange (nt ).reshape (nt , 1 , 1 )
320+
321+ lhs = v .backward if backward else v .forward
322+ eqn = Eq (lhs , v + u )
323+
324+ op0 = Operator (eqn , opt = 'noop' , name = 'op0' )
325+ op1 = Operator (eqn , opt = ('buffering' ,
326+ {'buf-async-degree' : async_degree }), name = 'op1' )
327+
328+ buffers = [i for i in FindSymbols ().visit (op1 )
329+ if i .is_Array and i ._mem_heap ]
330+ assert len (buffers ) == int (async_degree != 0 )
331+ if async_degree :
332+ assert buffers [0 ].symbolic_shape [0 ] == async_degree
333+
334+ for op in [op0 , op1 ]:
335+ args = op .arguments ()
336+ assert (args ['time_m' ], args ['time_M' ]) == expected_bounds
337+
338+ # Default bounds, either endpoint, a partial ring, and an empty interval
339+ time_m , time_M = expected_bounds
340+ for kwargs in [{}, {'time_m' : time_m , 'time_M' : time_m },
341+ {'time_m' : time_M , 'time_M' : time_M },
342+ {'time_m' : 3 , 'time_M' : 4 }, {'time_m' : 1 , 'time_M' : 0 }]:
343+ v .data [:] = 0
344+ v1 .data [:] = 0
345+ op0 .apply (** kwargs )
346+ op1 .apply (v = v1 , ** kwargs )
347+
348+ assert np .all (v .data == v1 .data )
349+
350+
351+ @pytest .mark .parametrize ('async_degree' , [None , 4 ])
352+ def test_two_homogeneous_buffers (async_degree ):
299353 nt = 10
300354 grid = Grid (shape = (4 , 4 ))
301355
@@ -308,8 +362,10 @@ def test_two_homogeneous_buffers():
308362 Eq (v .forward , u + v + u .backward + v .backward + 1. )]
309363
310364 op0 = Operator (eqns , opt = 'noop' )
311- op1 = Operator (eqns , opt = 'buffering' )
312- op2 = Operator (eqns , opt = ('buffering' , 'fuse' ))
365+ op1 = Operator (eqns , opt = ('buffering' ,
366+ {'buf-async-degree' : async_degree }))
367+ op2 = Operator (eqns , opt = ('buffering' , 'fuse' ,
368+ {'buf-async-degree' : async_degree }))
313369
314370 # Check generated code
315371 assert len (retrieve_iteration_tree (op1 )) == 5
@@ -323,8 +379,16 @@ def test_two_homogeneous_buffers():
323379 assert np .all (u .data == u1 .data )
324380 assert np .all (v .data == v1 .data )
325381
382+ u1 .data [:] = 0
383+ v1 .data [:] = 0
384+ op2 .apply (time_M = nt - 2 , u = u1 , v = v1 )
385+
386+ assert np .all (u .data == u1 .data )
387+ assert np .all (v .data == v1 .data )
388+
326389
327- def test_two_heterogeneous_buffers ():
390+ @pytest .mark .parametrize ('async_degree' , [None , 4 ])
391+ def test_two_heterogeneous_buffers (async_degree ):
328392 nt = 10
329393 grid = Grid (shape = (4 , 4 ))
330394
@@ -341,7 +405,8 @@ def test_two_heterogeneous_buffers():
341405 Eq (v .forward , u + v + v .backward )]
342406
343407 op0 = Operator (eqns , opt = 'noop' )
344- op1 = Operator (eqns , opt = 'buffering' )
408+ op1 = Operator (eqns , opt = ('buffering' ,
409+ {'buf-async-degree' : async_degree }))
345410
346411 # Check generated code
347412 assert len (retrieve_iteration_tree (op1 )) == 5
0 commit comments