Skip to content

Commit 959479e

Browse files
Merge pull request #961 from maelso/fix-redeclaration
Scheduler: Fix redeclaration
2 parents 4692baa + 185f894 commit 959479e

3 files changed

Lines changed: 54 additions & 16 deletions

File tree

‎devito/ir/iet/scheduler.py‎

Lines changed: 15 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
ExpressionBundle, Transformer, FindNodes, FindSymbols,
88
MapExprStmts, XSubs, iet_analyze)
99
from devito.symbolics import IntDiv, ccode, xreplace_indices
10-
from devito.tools import as_mapper, as_tuple
10+
from devito.tools import as_mapper, as_tuple, flatten
1111
from devito.types import ConditionalDimension
1212

1313
__all__ = ['iet_build', 'iet_insert_decls', 'iet_insert_casts']
@@ -168,7 +168,7 @@ def iet_insert_decls(iet, external):
168168
continue
169169
elif i._mem_stack:
170170
# On the stack
171-
allocator.push_object_on_stack(iet[0], i)
171+
allocator.push_array_on_stack(iet[0], i)
172172
else:
173173
# On the heap
174174
allocator.push_array_on_heap(i)
@@ -199,16 +199,21 @@ def __init__(self):
199199
self.stack = OrderedDict()
200200

201201
def push_object_on_stack(self, scope, obj):
202-
"""Define an Array or a composite type (e.g., a struct) on the stack."""
202+
"""Define a LocalObject on the stack."""
203203
handle = self.stack.setdefault(scope, OrderedDict())
204+
handle[obj] = Element(c.Value(obj._C_typename, obj.name))
204205

205-
if obj.is_LocalObject:
206-
handle[obj] = Element(c.Value(obj._C_typename, obj.name))
207-
else:
208-
shape = "".join("[%s]" % ccode(i) for i in obj.symbolic_shape)
209-
alignment = "__attribute__((aligned(%d)))" % obj._data_alignment
210-
value = "%s%s %s" % (obj.name, shape, alignment)
211-
handle[obj] = Element(c.POD(obj.dtype, value))
206+
def push_array_on_stack(self, scope, obj):
207+
"""Define an Array on the stack."""
208+
handle = self.stack.setdefault(scope, OrderedDict())
209+
210+
if obj in flatten(self.stack.values()):
211+
return
212+
213+
shape = "".join("[%s]" % ccode(i) for i in obj.symbolic_shape)
214+
alignment = "__attribute__((aligned(%d)))" % obj._data_alignment
215+
value = "%s%s %s" % (obj.name, shape, alignment)
216+
handle[obj] = Element(c.POD(obj.dtype, value))
212217

213218
def push_scalar_on_stack(self, scope, expr):
214219
"""Define a Scalar on the stack."""

‎tests/test_operator.py‎

Lines changed: 20 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,10 +6,12 @@
66
SparseFunction, SparseTimeFunction, Dimension, error, SpaceDimension,
77
NODE, CELL, dimensions, configuration, TensorFunction,
88
TensorTimeFunction, VectorFunction, VectorTimeFunction)
9-
from devito.ir.iet import (Expression, Iteration, FindNodes, IsPerfectIteration,
9+
from devito.ir.equations import ClusterizedEq
10+
from devito.ir.iet import (Conditional, Expression, Iteration, FindNodes,
11+
IsPerfectIteration, derive_parameters, iet_insert_decls,
1012
retrieve_iteration_tree)
1113
from devito.ir.support import Any, Backward, Forward
12-
from devito.symbolics import indexify, retrieve_indexed
14+
from devito.symbolics import ListInitializer, indexify, retrieve_indexed
1315
from devito.tools import flatten
1416
from devito.types import Array, Scalar
1517

@@ -1143,6 +1145,22 @@ def test_stack_vector_temporaries(self):
11431145
timers->section0 += (double)(end_section0.tv_sec-start_section0.tv_sec)\
11441146
+(double)(end_section0.tv_usec-start_section0.tv_usec)/1000000;""" in str(operator)
11451147

1148+
def test_conditional_declarations(self):
1149+
a = Array(name='a', dimensions=(x,), dtype=np.int32, scope='stack')
1150+
list_initialize = Expression(ClusterizedEq(Eq(a, ListInitializer([0, 0]))))
1151+
iet = Conditional(x < 3, list_initialize, list_initialize)
1152+
parameters = derive_parameters(iet, True)
1153+
iet = iet_insert_decls(iet, parameters)
1154+
assert str(iet[0]) == """\
1155+
if (x < 3)
1156+
{
1157+
int a[x_size] = {0, 0};
1158+
}
1159+
else
1160+
{
1161+
int a[x_size] = {0, 0};
1162+
}"""
1163+
11461164

11471165
class TestLoopScheduling(object):
11481166

‎tests/test_ops.py‎

Lines changed: 19 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13,13 +13,15 @@
1313
# thus invalidating all of the future tests. This is guaranteed by the
1414
# `pytestmark` above
1515
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
1618
from devito.ops.node_factory import OPSNodeFactory # noqa
1719
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
1921
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
2123
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
2325

2426

2527
class TestOPSExpression(object):
@@ -272,11 +274,24 @@ def test_create_ops_block(self, equation, expected):
272274
])
273275
def test_upper_bound(self, equation, expected):
274276
grid = Grid((5, 5))
275-
u = TimeFunction(name='u', grid=grid) # noqa
277+
u = TimeFunction(name='u', grid=grid) # noqa
276278
op = Operator(eval(equation))
277279

278280
assert expected in str(op.ccode)
279281

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+
280295
@pytest.mark.parametrize('equation,expected', [
281296
('Eq(u_2d.forward, u_2d + 1)',
282297
'[\'ops_dat_fetch_data(u_dat[(time_M)%(2)],0,&(u[(time_M)%(2)]));\','

0 commit comments

Comments
 (0)