|
4 | 4 |
|
5 | 5 | from conftest import skipif |
6 | 6 | from devito import ( |
7 | | - CondEq, ConditionalDimension, Constant, Dimension, Eq, Function, Grid, Operator, |
8 | | - SparseTimeFunction, SubDimension, SubDomain, TimeFunction, configuration, switchconfig |
| 7 | + CondEq, ConditionalDimension, Constant, CustomDimension, Dimension, Eq, Function, |
| 8 | + Grid, Operator, SparseTimeFunction, SubDimension, SubDomain, TimeFunction, |
| 9 | + configuration, switchconfig |
9 | 10 | ) |
10 | 11 | from devito.arch.archinfo import AppleArm |
11 | 12 | from devito.exceptions import CompilationError |
12 | | -from devito.ir import FindSymbols, retrieve_iteration_tree |
| 13 | +from devito.ir import ( |
| 14 | + Cluster, FindSymbols, Interval, IterationSpace, lower_exprs, retrieve_iteration_tree |
| 15 | +) |
| 16 | +from devito.passes.clusters.buffering import BufferDimension, expand_halo_transfers |
| 17 | +from devito.types import Array |
13 | 18 |
|
14 | 19 |
|
15 | 20 | def test_read_write(): |
@@ -64,6 +69,100 @@ def test_write_only(): |
64 | 69 | assert np.all(v.data == v1.data) |
65 | 70 |
|
66 | 71 |
|
| 72 | +@pytest.mark.parametrize('forward', [False, True]) |
| 73 | +def test_write_only_with_halo_source(forward): |
| 74 | + """ |
| 75 | + A buffered save of a Function with a populated halo must preserve that halo. |
| 76 | + """ |
| 77 | + nt = 5 |
| 78 | + grid = Grid(shape=(17, 17)) |
| 79 | + y = grid.dimensions[-1] |
| 80 | + |
| 81 | + u = TimeFunction(name='u', grid=grid, space_order=8) |
| 82 | + usave = TimeFunction(name='usave', grid=grid, space_order=8, save=nt) |
| 83 | + |
| 84 | + k = CustomDimension(name='k', parent=y, symbolic_min=1, |
| 85 | + symbolic_max=4, symbolic_size=4) |
| 86 | + |
| 87 | + eqns = [Eq(u.forward, u + 1), |
| 88 | + Eq(u.forward._subs(y, -k), -u.forward._subs(y, k)), |
| 89 | + Eq(usave, u.forward if forward else u)] |
| 90 | + |
| 91 | + op = Operator(eqns, opt='buffering', name='save_halo') |
| 92 | + op.apply(time_M=nt-2) |
| 93 | + |
| 94 | + hx, hy = usave._size_halo.left[1:] |
| 95 | + for t in range(nt-1): |
| 96 | + assert np.all(usave.data[t] == t + forward) |
| 97 | + actual = usave.data_with_halo[t, hx:hx + grid.shape[0], hy-4:hy] |
| 98 | + assert np.all(actual == -(t + forward)) |
| 99 | + |
| 100 | + |
| 101 | +@pytest.mark.parametrize('space_order, shift', [(0, 0), (8, -1), (8, 1), (10, 1)]) |
| 102 | +@switchconfig(autopadding=False) |
| 103 | +def test_write_only_with_halo_source_bounds(space_order, shift): |
| 104 | + grid = Grid(shape=(17, 17)) |
| 105 | + y = grid.dimensions[-1] |
| 106 | + |
| 107 | + u = TimeFunction(name='u', grid=grid, space_order=8) |
| 108 | + v = TimeFunction(name='v', grid=grid, space_order=space_order) |
| 109 | + usave = TimeFunction(name='usave', grid=grid, space_order=8, save=5) |
| 110 | + |
| 111 | + k = CustomDimension(name='k', parent=y, symbolic_min=1, |
| 112 | + symbolic_max=4, symbolic_size=4) |
| 113 | + |
| 114 | + eqns = [Eq(u.forward, u + 1), |
| 115 | + Eq(u.forward._subs(y, -k), -u.forward._subs(y, k)), |
| 116 | + Eq(usave, u.forward + v.forward._subs(y, y + shift))] |
| 117 | + |
| 118 | + if space_order == 10: |
| 119 | + # A wider halo accommodates the shifted read |
| 120 | + v.data_with_halo[:] = 2 |
| 121 | + op = Operator(eqns, opt='buffering', name='save_shifted_halo') |
| 122 | + op.apply(time_M=3) |
| 123 | + assert np.all(usave.data[3] == 6) |
| 124 | + hx, hy = usave._size_halo.left[1:] |
| 125 | + assert np.all(usave.data_with_halo[3, hx:hx + grid.shape[0], hy-4:hy] == -2) |
| 126 | + else: |
| 127 | + with pytest.raises(CompilationError, match='Insufficient halo for `v`'): |
| 128 | + Operator(eqns, opt='buffering') |
| 129 | + |
| 130 | + |
| 131 | +@pytest.mark.parametrize('mixed', [False, True]) |
| 132 | +def test_halo_transfers_non_time_dimension(mixed): |
| 133 | + s = Dimension(name='s') |
| 134 | + x = Dimension(name='x') |
| 135 | + u = Function(name='u', dimensions=(s, x), shape=(5, 17), |
| 136 | + halo=((0, 0), (4, 4))) |
| 137 | + usave = Function(name='usave', dimensions=(s, x), shape=(5, 17), |
| 138 | + halo=u.halo) |
| 139 | + db = BufferDimension('db', 0, 0, 1, s) |
| 140 | + b = Array(name='b', dimensions=(db, x), halo=usave.halo) |
| 141 | + k = CustomDimension(name='k', parent=x, symbolic_min=1, |
| 142 | + symbolic_max=4, symbolic_size=4) |
| 143 | + |
| 144 | + mirror = Cluster(lower_exprs(Eq(u[s+1, -k], -u[s+1, k])), |
| 145 | + IterationSpace([Interval(s), Interval(k)])) |
| 146 | + eqns = [Eq(usave[s, x], u[s+1, x])] |
| 147 | + if mixed: |
| 148 | + eqns.append(Eq(u[s, x], 0)) |
| 149 | + save = Cluster(lower_exprs(eqns), IterationSpace([Interval(s), Interval(x)])) |
| 150 | + mapper = {(usave, save.guards): b} |
| 151 | + |
| 152 | + if mixed: |
| 153 | + with pytest.raises(CompilationError, match='mixed Cluster'): |
| 154 | + expand_halo_transfers([mirror, save], mapper) |
| 155 | + return |
| 156 | + |
| 157 | + clusters = expand_halo_transfers([mirror, save], mapper) |
| 158 | + |
| 159 | + assert clusters[0] is mirror |
| 160 | + assert clusters[1].ispace[x].offsets == (-4, 4) |
| 161 | + # The streaming axis is not part of the halo footprint, even with a shifted read |
| 162 | + assert clusters[1].ispace[s] == save.ispace[s] |
| 163 | + assert clusters[1].exprs[0].args == save.exprs[0].args |
| 164 | + |
| 165 | + |
67 | 166 | def test_read_only(): |
68 | 167 | nt = 10 |
69 | 168 | grid = Grid(shape=(2, 2)) |
|
0 commit comments