|
1 | 1 | import numpy as np |
2 | 2 | import pytest |
3 | | -from sympy import Float, Symbol, diff, simplify, sympify |
| 3 | +from sympy import Float, S, Symbol, diff, simplify, sympify |
4 | 4 |
|
5 | 5 | from conftest import assert_structure |
6 | 6 | from devito import ( |
|
10 | 10 | ) |
11 | 11 | from devito.finite_differences import Derivative, Differentiable, diffify |
12 | 12 | from devito.finite_differences.differentiable import ( |
13 | | - Add, DiffDerivative, EvalDerivative, IndexDerivative, IndexSum, Weights, interp_for_fd |
| 13 | + Add, DiffDerivative, EvalDerivative, IndexDerivative, IndexSum, Weights, diff2sympy, |
| 14 | + interp_for_fd |
14 | 15 | ) |
15 | | -from devito.symbolics import indexify, retrieve_indexed |
| 16 | +from devito.finite_differences.tools import generate_indices |
| 17 | +from devito.ir.equations.algorithms import lower_exprs |
| 18 | +from devito.symbolics import indexify, retrieve_indexed, search, uxreplace |
16 | 19 | from devito.types.dimension import StencilDimension |
17 | 20 | from devito.warnings import DevitoWarning |
18 | 21 |
|
@@ -1093,6 +1096,161 @@ def test_index_derivative(self): |
1093 | 1096 |
|
1094 | 1097 | assert IndexDerivative(vi0*w, {x: i}) == vi1 |
1095 | 1098 |
|
| 1099 | + @pytest.mark.parametrize('staggered', [None, x, y]) |
| 1100 | + @pytest.mark.parametrize('deriv_order,fd_order', [ |
| 1101 | + (0, 4), (0, 16), (1, 2), (1, 4), (1, 16), (2, 4), (2, 16) |
| 1102 | + ]) |
| 1103 | + @pytest.mark.parametrize('offset,staggering', [ |
| 1104 | + (0, 'centered'), (S.Half, 'staggered'), (-S.Half, 'staggered'), |
| 1105 | + (1, 'centered'), (S.One/4, None) |
| 1106 | + ]) |
| 1107 | + def test_index_derivative_staggering(self, staggered, deriv_order, fd_order, |
| 1108 | + offset, staggering): |
| 1109 | + grid = Grid(shape=(10, 10)) |
| 1110 | + x, y = grid.dimensions |
| 1111 | + staggered = staggered(grid) if staggered else NODE |
| 1112 | + f = Function(name='f', grid=grid, space_order=16, staggered=staggered) |
| 1113 | + x0 = {x: f.indices_ref[x] + offset*x.spacing} |
| 1114 | + deriv = f.diff(x, deriv_order=deriv_order, fd_order=fd_order, x0=x0) |
| 1115 | + |
| 1116 | + evaluated = deriv._evaluate(expand=False) |
| 1117 | + if deriv_order == 0 and offset == 0: |
| 1118 | + assert evaluated == f |
| 1119 | + return |
| 1120 | + |
| 1121 | + lowered = lower_exprs(diff2sympy(evaluated)) |
| 1122 | + for i in (evaluated, lowered): |
| 1123 | + assert isinstance(i, IndexDerivative) |
| 1124 | + assert i.staggering == staggering |
| 1125 | + assert i.deriv_order == deriv_order |
| 1126 | + |
| 1127 | + # Classifying the staggering must not change the discrete stencil |
| 1128 | + sd, = evaluated.dimensions |
| 1129 | + terms = [w*evaluated.base.subs(sd, i) for w, i in |
| 1130 | + zip(evaluated.weights.function.weights, sd.range, strict=True)] |
| 1131 | + assert simplify(sum(terms) - deriv.evaluate) == 0 |
| 1132 | + |
| 1133 | + @pytest.mark.parametrize('offsets,staggerings,fd_order,weights', [ |
| 1134 | + ((0, S.Half), ('centered', 'staggered'), 2, None), |
| 1135 | + ((S.Half, S.One/4), ('staggered', None), 4, |
| 1136 | + [S.One/24, -S(9)/8, S(9)/8, -S.One/24]) |
| 1137 | + ]) |
| 1138 | + @pytest.mark.parametrize('lowered', [False, True]) |
| 1139 | + def test_index_derivative_staggering_identity(self, offsets, staggerings, fd_order, |
| 1140 | + weights, lowered): |
| 1141 | + grid = Grid(shape=(10,)) |
| 1142 | + x, = grid.dimensions |
| 1143 | + f = Function(name='f', grid=grid, space_order=4) |
| 1144 | + derivs = [f.dx(fd_order=fd_order, x0={x: x + i*x.spacing}, weights=weights) |
| 1145 | + for i in offsets] |
| 1146 | + derivs = [i._evaluate(expand=False) for i in derivs] |
| 1147 | + if lowered: |
| 1148 | + derivs = [lower_exprs(diff2sympy(i)) for i in derivs] |
| 1149 | + a, b = derivs |
| 1150 | + |
| 1151 | + # Distinct staggerings can generate exactly the same discrete stencil |
| 1152 | + assert tuple(i.staggering for i in derivs) == staggerings |
| 1153 | + assert a.args == b.args |
| 1154 | + assert a.mapper == b.mapper |
| 1155 | + assert a.weights.function.weights == b.weights.function.weights |
| 1156 | + assert a != b |
| 1157 | + assert len({a, b}) == 2 |
| 1158 | + assert a.compare(b) == -b.compare(a) != 0 |
| 1159 | + assert set((a + b).args) == {a, b} |
| 1160 | + assert a.evaluate == b.evaluate |
| 1161 | + |
| 1162 | + @pytest.mark.parametrize('offset,staggering', [ |
| 1163 | + (-S.Half, 'staggered'), (0, 'centered') |
| 1164 | + ]) |
| 1165 | + @pytest.mark.parametrize('side', [None, centered, left, right]) |
| 1166 | + @pytest.mark.parametrize('transpose', [False, True]) |
| 1167 | + def test_index_derivative_staggering_transpose(self, offset, staggering, side, |
| 1168 | + transpose): |
| 1169 | + grid = Grid(shape=(10,)) |
| 1170 | + x, = grid.dimensions |
| 1171 | + f = Function(name='f', grid=grid, space_order=4, staggered=x) |
| 1172 | + deriv = f.dx(x0={x: f.indices_ref[x] + offset*x.spacing}, side=side) |
| 1173 | + if transpose: |
| 1174 | + deriv = deriv.T |
| 1175 | + evaluated = deriv._evaluate(expand=False) |
| 1176 | + assert evaluated.staggering == staggering |
| 1177 | + assert lower_exprs(diff2sympy(evaluated)).staggering == staggering |
| 1178 | + assert simplify(evaluated.evaluate - deriv.evaluate) == 0 |
| 1179 | + |
| 1180 | + @pytest.mark.parametrize('dim', [x, y]) |
| 1181 | + @pytest.mark.parametrize('offset,staggering', [ |
| 1182 | + (-S.Half, 'staggered'), (0, 'centered') |
| 1183 | + ]) |
| 1184 | + def test_index_derivative_staggering_nested(self, dim, offset, staggering): |
| 1185 | + grid = Grid(shape=(10, 10)) |
| 1186 | + x, y = grid.dimensions |
| 1187 | + d = dim(grid) |
| 1188 | + f = Function(name='f', grid=grid, space_order=4) |
| 1189 | + deriv = f.dx(x0={x: x + x.spacing/2}) |
| 1190 | + deriv = deriv.diff(d, deriv_order=2, x0={d: d + offset*d.spacing}) |
| 1191 | + |
| 1192 | + evaluated = deriv._evaluate(expand=False) |
| 1193 | + lowered = lower_exprs(diff2sympy(evaluated)) |
| 1194 | + for expr in (evaluated, lowered): |
| 1195 | + iderivs = search(expr, IndexDerivative) |
| 1196 | + assert len(iderivs) == 2 |
| 1197 | + assert {(i.deriv_order, i.staggering) for i in iderivs} == { |
| 1198 | + (1, 'staggered'), (2, staggering) |
| 1199 | + } |
| 1200 | + |
| 1201 | + @pytest.mark.parametrize('lowered', [False, True]) |
| 1202 | + def test_index_derivative_staggering_rebuild(self, lowered): |
| 1203 | + grid = Grid(shape=(10,)) |
| 1204 | + x, = grid.dimensions |
| 1205 | + f = Function(name='f', grid=grid, space_order=4) |
| 1206 | + g = Function(name='g', grid=grid, space_order=4) |
| 1207 | + expr = f.dx(x0={x: x + x.spacing/2})._evaluate(expand=False) |
| 1208 | + sd, = expr.dimensions |
| 1209 | + base = g.subs(x, x + sd*x.spacing) |
| 1210 | + if lowered: |
| 1211 | + expr = lower_exprs(diff2sympy(expr)) |
| 1212 | + base = lower_exprs(base) |
| 1213 | + |
| 1214 | + for rebuilt in (expr.func(base*expr.weights), expr.subs(expr.base, base), |
| 1215 | + expr.xreplace({expr.base: base}), |
| 1216 | + uxreplace(expr, {expr.base: base})): |
| 1217 | + assert rebuilt.base == base |
| 1218 | + assert rebuilt.staggering == 'staggered' |
| 1219 | + assert rebuilt.deriv_order == 1 |
| 1220 | + |
| 1221 | + # Missing metadata must stay distinct from a known centered interpolation |
| 1222 | + unknown = expr._rebuild(staggering=None, deriv_order=None) |
| 1223 | + on_grid = expr._rebuild(staggering='centered', deriv_order=0) |
| 1224 | + assert unknown.staggering is None |
| 1225 | + assert len({expr, unknown, on_grid}) == 3 |
| 1226 | + assert unknown.compare(on_grid) == -on_grid.compare(unknown) != 0 |
| 1227 | + assert set((expr + unknown + on_grid).args) == {expr, unknown, on_grid} |
| 1228 | + |
| 1229 | + def test_index_derivative_staggering_unknown(self): |
| 1230 | + grid = Grid(shape=(10, 10)) |
| 1231 | + x, y = grid.dimensions |
| 1232 | + f = Function(name='f', grid=grid, space_order=4) |
| 1233 | + exprs = (f.dx45, f.dx(x0={x: 1})) |
| 1234 | + for expr in exprs: |
| 1235 | + lowered = lower_exprs(diff2sympy(expr._evaluate(expand=False))) |
| 1236 | + iderivs = search(lowered, IndexDerivative) |
| 1237 | + assert iderivs |
| 1238 | + assert all(i.staggering is None for i in iderivs) |
| 1239 | + |
| 1240 | + @pytest.mark.parametrize('offset,staggering', [ |
| 1241 | + (0, 'centered'), (1.0, 'centered'), (-1.0, 'centered'), |
| 1242 | + (0.5, 'staggered'), (-1.5, 'staggered'), (0.25, None), (0.50000001, None) |
| 1243 | + ]) |
| 1244 | + def test_index_set_staggering(self, offset, staggering): |
| 1245 | + grid = Grid(shape=(10,)) |
| 1246 | + x, = grid.dimensions |
| 1247 | + f = Function(name='f', grid=grid, space_order=4, staggered=x) |
| 1248 | + indices, _ = generate_indices(f, x, 4, |
| 1249 | + x0={x: f.indices_ref[x] + offset*x.spacing}) |
| 1250 | + for i in (indices, indices.transpose(), indices.shift(-x.spacing/2), |
| 1251 | + indices.transpose().shift(-x.spacing/2)): |
| 1252 | + assert i.staggering == staggering |
| 1253 | + |
1096 | 1254 | def test_dx2(self): |
1097 | 1255 | grid = Grid(shape=(4, 4)) |
1098 | 1256 |
|
|
0 commit comments