Skip to content
12 changes: 10 additions & 2 deletions loki/expression/mappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,10 +226,14 @@ class LokiWalkMapper(WalkMapper):
"""
# pylint: disable=abstract-method

def __init__(self, recurse_var_parent=True, **kwargs):
super().__init__(**kwargs)
self.recurse_var_parent = recurse_var_parent

def map_variable_symbol(self, expr, *args, **kwargs):
if not self.visit(expr):
return
if expr.parent:
if expr.parent and self.recurse_var_parent:
self.rec(expr.parent, *args, **kwargs)
self.post_visit(expr, *args, **kwargs)

Expand Down Expand Up @@ -745,8 +749,12 @@ class SubstituteExpressionsMapper(LokiIdentityMapper):

def __init__(self, expr_map):
super().__init__()
from loki.expression.symbols import MetaSymbol # pylint: disable=import-outside-toplevel,cyclic-import

self.expr_map = expr_map
self.expr_map = dict(expr_map)
for expr, replacement in list(expr_map.items()):
if isinstance(expr, MetaSymbol) and expr._symbol not in self.expr_map:
self.expr_map[expr._symbol] = replacement
for expr in self.expr_map.keys():
setattr(self, expr.mapper_method, self.map_from_expr_map)

Expand Down
24 changes: 23 additions & 1 deletion loki/expression/tests/test_expression.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
available_frontends, OMNI, HAVE_FP, parse_fparser_expression
)
from loki.ir import (
nodes as ir, FindNodes, FindVariables, FindExpressions,
nodes as ir, FindNodes, FindUsedVariables, FindVariables, FindExpressions,
FindInlineCalls, SubstituteExpressions
)
from loki.tools import (
Expand Down Expand Up @@ -1379,6 +1379,28 @@ def test_nested_derived_type_substitution():
assert fgen(new_expr) == 'ydml_phy_mf%yrphy3%n_spband'


def test_find_used_variables_skips_parent_components():
"""Collect only the used member expression rather than its parent chain."""
expr = sym.Array(
name='field',
parent=sym.Scalar(name='geom', parent=sym.Scalar(name='state')),
dimensions=(sym.Scalar(name='jk'),)
)

all_variables = FindVariables(unique=False).visit(expr)
used_variables = FindUsedVariables(unique=False).visit(expr)

assert any(var == 'state%geom%field(jk)' for var in all_variables)
assert any(var == 'state%geom' for var in all_variables)
assert any(var == 'state' for var in all_variables)
assert any(var == 'jk' for var in all_variables)

assert any(var == 'state%geom%field(jk)' for var in used_variables)
assert any(var == 'jk' for var in used_variables)
assert not any(var == 'state%geom' for var in used_variables)
assert not any(var == 'state' for var in used_variables)


@pytest.mark.parametrize('frontend', available_frontends())
def test_variable_in_declaration_initializer(frontend):
"""
Expand Down
14 changes: 13 additions & 1 deletion loki/ir/expr_visitors.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
)

__all__ = [
'ExpressionFinder', 'FindExpressions', 'FindVariables',
'ExpressionFinder', 'FindExpressions', 'FindVariables', 'FindUsedVariables',
'FindTypedSymbols', 'FindInlineCalls', 'FindLiterals',
'FindRealLiterals', 'ExpressionTransformer',
'SubstituteExpressions', 'SubstituteStringExpressions',
Expand Down Expand Up @@ -184,6 +184,18 @@ class FindVariables(ExpressionFinder):
retriever = ExpressionRetriever(lambda e: isinstance(e, (Scalar, Array, DeferredTypeSymbol)))


class FindUsedVariables(ExpressionFinder):
"""
A visitor to collect variables without recursing into parent components
of derived-type member accesses.

See :class:`ExpressionFinder` for further details.
"""
retriever = ExpressionRetriever(
lambda e: isinstance(e, (Scalar, Array, DeferredTypeSymbol)), recurse_var_parent=False
)


class FindInlineCalls(ExpressionFinder):
"""
A visitor to collect all :any:`InlineCall` symbols used in an IR tree.
Expand Down
1 change: 1 addition & 0 deletions loki/transformations/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
from loki.transformations.idempotence import * # noqa
from loki.transformations.inline import * # noqa
from loki.transformations.parametrise import * # noqa
from loki.transformations.replace_kernel import * # noqa
from loki.transformations.remove_code import * # noqa
from loki.transformations.sanitise import * # noqa
from loki.transformations.single_column import * # noqa
Expand Down
12 changes: 9 additions & 3 deletions loki/transformations/array_indexing/promote.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
]


def promote_variables(routine, variable_names, pos, index=None, size=None):
def promote_variables(routine, variable_names, pos, index=None, size=None, ignore_index_undefined=False):
"""
Promote a list of variables by inserting new array dimensions of given size
and updating all uses of these variables with a given index expression.
Expand Down Expand Up @@ -56,6 +56,9 @@ def promote_variables(routine, variable_names, pos, index=None, size=None):
The size of the dimension (or tuple for multi-dimension promotion) to
insert at `pos`. When this is provided, the declaration of variables
is updated accordingly.
ignore_index_undefined : bool, optional
When `True`, keep the provided index expression even if dataflow
analysis cannot prove it is live at a particular use site.
"""
variable_names = {name.lower() for name in variable_names}

Expand All @@ -80,8 +83,11 @@ def promote_variables(routine, variable_names, pos, index=None, size=None):

# We use the given index expression in this node if all
# variables therein are defined, otherwise we use `:`
node_index = tuple(i if v <= node.live_symbols else sym.RangeIndex((None, None))
for i, v in zip(index, index_vars))
if not ignore_index_undefined:
node_index = tuple(i if v <= node.live_symbols else sym.RangeIndex((None, None))
for i, v in zip(index, index_vars))
else:
node_index = tuple(i for i, _ in zip(index, index_vars))

var_map = {}
for var in var_list:
Expand Down
28 changes: 28 additions & 0 deletions loki/transformations/array_indexing/tests/test_array_promote.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from loki.jit_build import jit_compile_and_run
from loki.expression import symbols as sym
from loki.frontend import available_frontends
from loki.ir import FindNodes, nodes as ir

from loki.transformations.array_indexing.promote import promote_variables

Expand Down Expand Up @@ -134,3 +135,30 @@ def test_transform_promote_variables(tmp_path, frontend):
assert scalar == n*(n+1)//2
assert np.all(vector[:-1] == np.array(list(range(n + 1, 2*n)), order='F', dtype=np.int32))
assert vector[-1] == 3*n


@pytest.mark.parametrize('frontend', available_frontends())
def test_promote_variables_keeps_explicit_index_when_requested(frontend):
"""Keep an unresolved promotion index verbatim when requested by the caller."""
fcode = """
subroutine promote_unknown_index(arr, n)
implicit none
integer, intent(in) :: n
integer, intent(inout) :: arr(n)
integer :: idx
integer :: tmp(n)

arr(:) = tmp(:)
end subroutine promote_unknown_index
""".strip()
routine = Subroutine.from_source(fcode, frontend=frontend)

promote_variables(
routine, ['tmp'], pos=-1, index=routine.variable_map['idx'], size=routine.variable_map['n'],
ignore_index_undefined=True
)

assign = FindNodes(ir.Assignment).visit(routine.body)[0]

assert routine.variable_map['tmp'].shape == (routine.variable_map['n'], routine.variable_map['n'])
assert assign.rhs == 'tmp(:, idx)'
Loading
Loading