|
13 | 13 | # thus invalidating all of the future tests. This is guaranteed by the |
14 | 14 | # `pytestmark` above |
15 | 15 | from devito import Eq, Function, Grid, Operator, TimeFunction, configuration # noqa |
| 16 | +from devito.ir.equations import ClusterizedEq # noqa |
| 17 | +from devito.ir.iet import Conditional, Expression, derive_parameters, iet_insert_decls # noqa |
16 | 18 | from devito.ops.node_factory import OPSNodeFactory # noqa |
17 | 19 | from devito.ops.transformer import create_ops_arg, create_ops_dat, make_ops_ast, to_ops_stencil # noqa |
18 | | -from devito.ops.types import OpsAccessible, OpsDat, OpsStencil, OpsBlock # noqa |
| 20 | +from devito.ops.types import Array, OpsAccessible, OpsDat, OpsStencil, OpsBlock # noqa |
19 | 21 | from devito.ops.utils import namespace, AccessibleInfo, OpsDatDecl, OpsArgDecl # noqa |
20 | | -from devito.symbolics import Byref, Literal, indexify # noqa |
| 22 | +from devito.symbolics import Byref, ListInitializer, Literal, indexify # noqa |
21 | 23 | from devito.tools import dtype_to_cstr # noqa |
22 | | -from devito.types import Buffer, Constant, Symbol # noqa |
| 24 | +from devito.types import Buffer, Constant, DefaultDimension, Symbol # noqa |
23 | 25 |
|
24 | 26 |
|
25 | 27 | class TestOPSExpression(object): |
@@ -272,11 +274,24 @@ def test_create_ops_block(self, equation, expected): |
272 | 274 | ]) |
273 | 275 | def test_upper_bound(self, equation, expected): |
274 | 276 | grid = Grid((5, 5)) |
275 | | - u = TimeFunction(name='u', grid=grid) # noqa |
| 277 | + u = TimeFunction(name='u', grid=grid) # noqa |
276 | 278 | op = Operator(eval(equation)) |
277 | 279 |
|
278 | 280 | assert expected in str(op.ccode) |
279 | 281 |
|
| 282 | + @pytest.mark.parametrize('equation, declaration', [ |
| 283 | + ('Eq(u.forward, u+1)', |
| 284 | + 'int OPS_Kernel_0_range[4]') |
| 285 | + ]) |
| 286 | + def test_single_declaration(self, equation, declaration): |
| 287 | + grid = Grid((5, 5)) |
| 288 | + u = TimeFunction(name='u', grid=grid) # noqa |
| 289 | + op = Operator(eval(equation)) |
| 290 | + |
| 291 | + occurrences = [i for i in str(op.ccode).split('\n') if declaration in i] |
| 292 | + |
| 293 | + assert len(occurrences) == 1 |
| 294 | + |
280 | 295 | @pytest.mark.parametrize('equation,expected', [ |
281 | 296 | ('Eq(u_2d.forward, u_2d + 1)', |
282 | 297 | '[\'ops_dat_fetch_data(u_dat[(time_M)%(2)],0,&(u[(time_M)%(2)]));\',' |
|
0 commit comments