Skip to content

Commit f976f24

Browse files
authored
Merge pull request #3036 from devitocodes/dummy-bug
compiler: Fix CIRE placeholder leaking into Operator parameters
2 parents 4f386a6 + d3a02fb commit f976f24

2 files changed

Lines changed: 30 additions & 15 deletions

File tree

‎devito/passes/clusters/aliases.py‎

Lines changed: 5 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -813,23 +813,13 @@ def make_variant(schedule, exprs, mapper):
813813
Create a Variant from a Schedule and the corresponding expressions.
814814
"""
815815
# Some aliases may have been discarded along the way, and for
816-
# them we reinstate the original sub-expressions
816+
# them we reinstate the original sub-expressions. `mapper.extracted` binds
817+
# each placeholder to exactly the (sub-)expression it replaced, including
818+
# for compound extractions, where the placeholder stands for a subset of
819+
# `expr.args` rather than for the whole `expr`
817820
retained = flatten(sa.aliaseds for sa in schedule)
818821

819-
subs = {}
820-
for k, v in mapper.items():
821-
if v in retained:
822-
continue
823-
elif isinstance(v, dict):
824-
# E.g., `mapper = {u[t0, x+3, y+3] + u[t0, x+3, y+4]:
825-
# {u[t0, x+3, y+4]: None, u[t0, x+3, y+3]: dummy0}}`
826-
try:
827-
v1, = [i for i in v.values() if i not in retained]
828-
except ValueError:
829-
continue
830-
subs[v1] = k
831-
else:
832-
subs[v] = k
822+
subs = {v: k for k, v in mapper.extracted.items() if v not in retained}
833823

834824
exprs = [uxreplace(e, subs) for e in exprs]
835825

‎tests/test_dse.py‎

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2748,6 +2748,31 @@ def test_split_cond(self):
27482748
scalars = [i for i in FindSymbols().visit(op) if isinstance(i, Temp)]
27492749
assert len(scalars) == 0
27502750

2751+
def test_split_cond_compound_alias(self):
2752+
"""
2753+
MFE for the case in which a CIRE placeholder Symbol survives into
2754+
`op.parameters` because the compound extraction it stands for is
2755+
discarded by `lower_aliases` (scalar alias + guarded Cluster).
2756+
"""
2757+
grid = Grid((11, 11))
2758+
time = grid.time_dim
2759+
2760+
u = TimeFunction(name='u', grid=grid, time_order=2, space_order=2)
2761+
u1 = TimeFunction(name='u1', grid=grid, time_order=2, space_order=2)
2762+
2763+
ct = ConditionalDimension(name='ct', parent=time, factor=2)
2764+
2765+
eqn = Eq(u.forward, u + sin(time)*cos(time), implicit_dims=ct)
2766+
2767+
op0 = Operator(eqn, opt='noop')
2768+
op1 = Operator(eqn, opt=('advanced', {'cire-mingain': 0}))
2769+
2770+
assert not any(i.name.startswith('dummy') for i in op1.parameters)
2771+
2772+
op0.apply(time_M=5)
2773+
op1.apply(time_M=5, u=u1)
2774+
assert np.allclose(u.data, u1.data, rtol=1e-5)
2775+
27512776
def test_split_cond_multi_alias(self):
27522777
grid = Grid((11, 11))
27532778
time = grid.time_dim

0 commit comments

Comments
 (0)