@@ -50,12 +50,18 @@ def test_scheduling_after_rewrite():
5050 op = Operator ([eqn1 ] + eqn2 + [eqn3 ])
5151 trees = retrieve_iteration_tree (op )
5252
53- # Check loop nest structure
54- assert all (
55- i .dim is j for i , j in zip (trees [0 ], grid .dimensions , strict = True )
56- ) # time invariant
57- assert trees [1 ].root .dim is grid .time_dim
58- assert all (trees [1 ].root .dim is tree .root .dim for tree in trees [1 :])
53+ # Check loop nest structure. `CireInvariantsSparse` hoists the sparse
54+ # position preamble as an extra time-invariant tree ahead of the time
55+ # loops, alongside the dense sin(const) preamble.
56+ invariants , in_time = [], []
57+ for t in trees :
58+ (in_time if t .root .dim is grid .time_dim else invariants ).append (t )
59+ assert invariants # sin(const) + the sparse-position preamble(s)
60+ assert in_time # the actual time-loop clusters
61+ # Invariants must all precede the time-loop trees
62+ for inv in invariants :
63+ for tt in in_time :
64+ assert trees .index (inv ) < trees .index (tt )
5965
6066
6167@pytest .mark .parametrize ('expr,expected' , [
@@ -1399,7 +1405,10 @@ def test_catch_duplicate_from_different_clusters(self):
13991405 op = Operator (eqns , opt = ('advanced' , {'cire-mingain' : 100 }))
14001406
14011407 arrays = [i for i in FindSymbols ().visit (op ) if i .is_Array ]
1402- assert len (arrays ) == 3
1408+ # 3 dense-derivative aliases + 4 sparse-position preambles (pos_x,
1409+ # pos_y and their int(floor(...)) tabulations) from
1410+ # `CireInvariantsSparse`.
1411+ assert len (arrays ) == 7
14031412 assert all (i ._mem_heap and not i ._mem_external for i in arrays )
14041413
14051414 def test_discarded_compound (self ):
@@ -1542,7 +1551,9 @@ def test_drop_redundants_after_fusion(self, rotate):
15421551 op = Operator (eqns , opt = ('advanced' , {'cire-rotate' : rotate }))
15431552
15441553 arrays = [i for i in FindSymbols ().visit (op ) if i .is_Array ]
1545- assert len (arrays ) == 2
1554+ # 2 dense aliases + 4 sparse-position preambles per sparse function
1555+ # (pos_x, pos_y, int_floor_x, int_floor_y) x 1 SparseTimeFunction.
1556+ assert len (arrays ) == 6
15461557 assert all (i ._mem_heap and not i ._mem_external for i in arrays )
15471558
15481559 def test_full_shape_big_temporaries (self ):
@@ -2960,10 +2971,13 @@ def test_fullopt(self):
29602971 assert summary0 [('section1' , None )].ops == 44
29612972 assert np .isclose (summary0 [('section0' , None )].oi , 3.136 , atol = 0.001 )
29622973
2963- assert summary1 [('section0' , None )].ops == 31
2964- assert summary1 [('section1' , None )].ops == 88
2965- assert summary1 [('section2' , None )].ops == 25
2966- assert np .isclose (summary1 [('section0' , None )].oi , 1.767 , atol = 0.001 )
2974+ # section0 is now the sparse-position preamble hoisted out of the
2975+ # time loop by `CireInvariantsSparse`; the dense stencil is section1.
2976+ assert summary1 [('section0' , None )].ops == 9
2977+ assert summary1 [('section1' , None )].ops == 31
2978+ assert summary1 [('section2' , None )].ops == 40
2979+ assert summary1 [('section3' , None )].ops == 25
2980+ assert np .isclose (summary1 [('section1' , None )].oi , 0.884 , atol = 0.001 )
29672981
29682982 assert np .allclose (u0 .data , u1 .data , atol = 10e-5 )
29692983 assert np .allclose (rec0 .data , rec1 .data , atol = 10e-5 )
@@ -3022,9 +3036,13 @@ def test_fullopt(self):
30223036 assert np .allclose (self .tti_noopt [0 ].data , v .data , atol = 10e-1 )
30233037 assert np .allclose (self .tti_noopt [1 ].data , rec .data , atol = 10e-1 )
30243038
3025- # Check expected opcount/oi
3026- assert summary [('section1' , None )].ops == 92
3027- assert np .isclose (summary [('section1' , None )].oi , 1.99 , atol = 0.001 )
3039+ # Check expected opcount/oi. The dense TTI stencil moved from
3040+ # `section1` to `section2` (inside multipass0) because sparse-position
3041+ # preambles now occupy the leading sections. Access the sops directly
3042+ # via the profiler's subsection map.
3043+ op_fwd = wavesolver .op_fwd ()
3044+ stencil = op_fwd ._profiler ._subsections ['multipass0' ]['section2' ]
3045+ assert stencil .sops == 92
30283046
30293047 # With optimizations enabled, there should be exactly four BlockDimensions
30303048 op = wavesolver .op_fwd ()
@@ -3035,15 +3053,15 @@ def test_fullopt(self):
30353053 assert y .parent is y0_blk0
30363054 assert not x ._defines & y ._defines
30373055
3038- # Also, in this operator, we expect six temporary Arrays:
3039- # * all of the six Arrays are allocated on the heap
3040- # * with OpenMP:
3041- # four Arrays are defined globally for the cos/sin temporaries
3042- # 3 Arrays are defined globally for the sparse positions temporaries
3043- # and two additional bock-sized Arrays are defined locally
3056+ # Temporary Arrays expected :
3057+ # * 4 global for the cos/sin temporaries
3058+ # * 6 global for the sparse-position preambles: per grid dim, one
3059+ # fp64 `pos` and one int32 `int(floor(pos))`, coalesced across
3060+ # the src+rec inject/interp.
3061+ # * 2 block-local
30443062 arrays = [i for i in FindSymbols ().visit (op ) if i .is_Array ]
30453063 extra_arrays = 2
3046- assert len (arrays ) == 4 + extra_arrays
3064+ assert len (arrays ) == 4 + 6 + extra_arrays
30473065 assert all (i ._mem_heap and not i ._mem_external for i in arrays )
30483066 bns , pbs = assert_blocking (op , {'x0_blk0' })
30493067
@@ -3076,19 +3094,26 @@ def test_fullopt_w_mpi(self, mode):
30763094 ])
30773095 def test_opcounts (self , space_order , expected ):
30783096 op = self .tti_operator (opt = 'advanced' , space_order = space_order )
3079- sections = list (op .op_fwd ()._profiler ._sections .values ())
3080- assert sections [1 ].sops == expected
3097+ # The dense TTI stencil moved inside `multipass0` when sparse-position
3098+ # preambles took the leading section slots.
3099+ op_fwd = op .op_fwd ()
3100+ stencil = op_fwd ._profiler ._subsections ['multipass0' ]['section2' ]
3101+ assert stencil .sops == expected
30813102
30823103 @switchconfig (profiling = 'advanced' )
30833104 @pytest .mark .parametrize ('space_order,exp_ops,exp_arrays' , [
3084- (4 , 122 , 6 ), (8 , 225 , 7 )
3105+ (4 , 122 , 12 ), (8 , 225 , 13 )
30853106 ])
30863107 def test_opcounts_adjoint (self , space_order , exp_ops , exp_arrays ):
30873108 wavesolver = self .tti_operator (space_order = space_order ,
30883109 opt = ('advanced' , {'openmp' : False }))
30893110 op = wavesolver .op_adj ()
30903111
3091- assert op ._profiler ._sections ['section1' ].sops == exp_ops
3112+ # Dense stencil is nested inside `multipass0` (sparse-position
3113+ # preambles occupy the leading sections). `exp_arrays` grows by 6 to
3114+ # include the sparse `pos`/`int_floor` preambles (2 per grid dim).
3115+ stencil = op ._profiler ._subsections ['multipass0' ]['section2' ]
3116+ assert stencil .sops == exp_ops
30923117 assert len ([i for i in FindSymbols ().visit (op ) if i .is_Array ]) == exp_arrays
30933118
30943119
0 commit comments