Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
97 changes: 59 additions & 38 deletions src/gt4py/next/iterator/transforms/inline_fundefs.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,76 +6,97 @@
# Please, refer to the LICENSE file in the root directory.
# SPDX-License-Identifier: BSD-3-Clause

from typing import Any, Dict
import dataclasses
import graphlib
from typing import Callable

from gt4py.eve import NodeTranslator, PreserveLocationVisitor
from gt4py.next.iterator import ir as itir
from gt4py.next.iterator.ir_utils import ir_makers as im
from gt4py.next.iterator.transforms import symbol_ref_utils


class InlineFundefs(PreserveLocationVisitor, NodeTranslator):
def visit_SymRef(self, node: itir.SymRef, *, symtable: Dict[str, Any]):
if node.id in symtable and isinstance(
(symbol := symtable[node.id]), itir.FunctionDefinition
):
return itir.Lambda(
params=self.generic_visit(symbol.params, symtable=symtable),
expr=self.generic_visit(symbol.expr, symtable=symtable),
)
return self.generic_visit(node)
def _sorted_by_dependency(
function_definitions: list[itir.FunctionDefinition],
) -> list[itir.FunctionDefinition]:
"""Order function definitions such that each one only references its predecessors."""
fundefs = {str(fundef.id): fundef for fundef in function_definitions}
dependencies = {
name: symbol_ref_utils.collect_symbol_refs(fundef.expr, fundefs.keys())
for name, fundef in fundefs.items()
}
return [fundefs[name] for name in graphlib.TopologicalSorter(dependencies).static_order()]


@dataclasses.dataclass(frozen=True)
class _WrapStatementExpressions(PreserveLocationVisitor, NodeTranslator):
"""Replace every expression a statement is made of by `wrap` applied to it."""

wrap: Callable[[itir.Expr], itir.Expr]

def visit_SetAt(self, node: itir.SetAt) -> itir.SetAt:
return itir.SetAt(expr=self.wrap(node.expr), domain=node.domain, target=node.target)

def visit_Program(self, node: itir.Program):
return self.generic_visit(node, symtable=node.annex.symtable)
def visit_IfStmt(self, node: itir.IfStmt) -> itir.IfStmt:
return itir.IfStmt(
cond=self.wrap(node.cond),
true_branch=self.visit(node.true_branch),
false_branch=self.visit(node.false_branch),
)


def prune_unreferenced_fundefs(program: itir.Program) -> itir.Program:
def inline_fundefs(program: itir.Program) -> itir.Program:
"""
Remove all function declarations that are never called.
Turn the function definitions of a program into `let` bindings of its statements.

Every statement expression is wrapped in a `let` binding all function definitions to the
corresponding lambdas. Since the function definitions thereby become regular bindings, the
usual scoping rules apply, in particular a lambda parameter of the same name shadows a
function definition instead of being erroneously replaced by it. Bindings that are unused,
e.g. because the function definition is never called, are removed by dead code elimination.

>>> from gt4py.next import common
>>> from gt4py.next.iterator.ir_utils import ir_makers as im
>>> fun1 = itir.FunctionDefinition(
... id="fun1",
... params=[im.sym("a")],
... expr=im.deref("a"),
... )
>>> fun2 = itir.FunctionDefinition(
... id="fun2",
... params=[im.sym("a")],
... expr=im.deref("a"),
... )
>>> fun1 = itir.FunctionDefinition(id="fun1", params=[im.sym("a")], expr=im.deref("a"))
>>> fun2 = itir.FunctionDefinition(id="fun2", params=[im.sym("a")], expr=im.call("fun1")("a"))
>>> IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL)
>>> program = itir.Program(
... id="testee",
... function_definitions=[fun1, fun2],
... function_definitions=[fun2, fun1],
... params=[im.sym("inp"), im.sym("out")],
... declarations=[],
... body=[
... itir.SetAt(
... expr=im.call("fun1")("inp"),
... expr=im.call("fun2")("inp"),
... domain=im.domain("cartesian_domain", {IDim: (0, 10)}),
... target=im.ref("out"),
... )
... ],
... )
>>> print(prune_unreferenced_fundefs(program))
>>> print(inline_fundefs(program))
testee(inp, out) {
fun1 = λ(a) → ·a;
out @ c⟨ IDimₕ: [0, 10[ ⟩ ← fun1(inp);
out @ c⟨ IDimₕ: [0, 10[ ⟩ ← (λ(fun1) → (λ(fun2) → fun2(inp))(λ(a) → fun1(a)))(λ(a) → ·a);
}
"""
fun_names = [fun.id for fun in program.function_definitions]
referenced_fun_names = symbol_ref_utils.collect_symbol_refs(program.body, fun_names)
if not program.function_definitions:
return program

# dependent function definitions are bound further inside, such that they see the function
# definitions they reference
bindings = [
(im.sym(fundef.id), im.lambda_(*fundef.params)(fundef.expr))
for fundef in _sorted_by_dependency(program.function_definitions)
]

new_fun_defs = []
for fun_def in program.function_definitions:
if fun_def.id in referenced_fun_names:
new_fun_defs.append(fun_def)
def wrap(expr: itir.Expr) -> itir.Expr:
for param, value in reversed(bindings):
expr = im.let(param, value)(expr)
return expr

return itir.Program(
id=program.id,
function_definitions=new_fun_defs,
function_definitions=[],
params=program.params,
declarations=program.declarations,
body=program.body,
body=_WrapStatementExpressions(wrap).visit(program.body),
)
7 changes: 2 additions & 5 deletions src/gt4py/next/iterator/transforms/pass_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,9 +160,7 @@ def apply_common_transforms(
uids = utils.IDGeneratorPool()

ir = MergeLet().visit(ir)
ir = inline_fundefs.InlineFundefs().visit(ir)

ir = inline_fundefs.prune_unreferenced_fundefs(ir)
ir = inline_fundefs.inline_fundefs(ir)
ir = NormalizeShifts().visit(ir)

# TODO(tehrengruber): Many iterator test contain lifts that need to be inlined, e.g.
Expand Down Expand Up @@ -284,8 +282,7 @@ def apply_fieldview_transforms(
ir, offset_provider, None, use_max_domain_range_on_unstructured_shift
)

ir = inline_fundefs.InlineFundefs().visit(ir)
ir = inline_fundefs.prune_unreferenced_fundefs(ir)
ir = inline_fundefs.inline_fundefs(ir)
# required for dead-code-elimination and `prune_empty_concat_where` pass
ir = concat_where.expand_tuple_args(ir, offset_provider_type=offset_provider_type) # type: ignore[assignment] # always an itir.Program
ir = expand_tuple_maps.ExpandTupleMaps.apply(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -45,3 +45,22 @@ def consume_constants(input: cases.IFloatField) -> cases.IFloatField:
cases.verify_with_default_data(
cartesian_case, consume_constants, ref=lambda input: Constants.PI * Constants.E * input
)


def test_param_shadowing_closure_var_field_operator(cartesian_case):
# `scale` is a closure variable of `testee` and therefore becomes a function definition of the
# lowered program. The parameter of `shadow` has the same name and must not be confused with it.
@gtx.field_operator
def scale(a: cases.IFloatField) -> cases.IFloatField:
return 2.0 * a

@gtx.field_operator
def shadow(scale: cases.IFloatField) -> cases.IFloatField:
return scale + 1.0

@gtx.program
def testee(a: cases.IFloatField, out: cases.IFloatField):
scale(a, out=out)
shadow(out, out=out)

cases.verify_with_default_data(cartesian_case, testee, ref=lambda a: 2.0 * a + 1.0)
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,8 @@ def laplacian(inp):
)


@fundef
# note: not a `fundef` as it would be polymorphic in `d`, which the type inference does not
# support; tracing inlines it as a lambda instead
def flux(d):
def flux_impl(inp):
lap = lift(laplacian)(inp)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
# GT4Py - GridTools Framework
#
# Copyright (c) 2014-2024, ETH Zurich
# All rights reserved.
#
# Please, refer to the LICENSE file in the root directory.
# SPDX-License-Identifier: BSD-3-Clause

from gt4py.next import common
from gt4py.next.iterator import ir as itir
from gt4py.next.iterator.ir_utils import ir_makers as im
from gt4py.next.iterator.transforms import inline_fundefs
from gt4py.next.type_system import type_specifications as ts


TDim = common.Dimension(value="TDim")
int_type = ts.ScalarType(kind=ts.ScalarKind.INT32)
domain = im.domain(common.GridType.CARTESIAN, {TDim: (0, 1)})


def program_factory(
body: list[itir.Stmt], function_definitions: list[itir.FunctionDefinition]
) -> itir.Program:
return itir.Program(
id="testee",
function_definitions=function_definitions,
params=[
im.sym("inp", ts.FieldType(dims=[TDim], dtype=int_type)),
im.sym("out", ts.FieldType(dims=[TDim], dtype=int_type)),
],
declarations=[],
body=body,
)


def test_simple():
fun = itir.FunctionDefinition(id="fun", params=[im.sym("a")], expr=im.deref("a"))
testee = program_factory(
body=[itir.SetAt(expr=im.call("fun")("inp"), target=im.ref("out"), domain=domain)],
function_definitions=[fun],
)
expected = program_factory(
body=[
itir.SetAt(
expr=im.let("fun", im.lambda_("a")(im.deref("a")))(im.call("fun")("inp")),
target=im.ref("out"),
domain=domain,
)
],
function_definitions=[],
)

actual = inline_fundefs.inline_fundefs(testee)
assert actual == expected


def test_unreferenced_fundef():
# unreferenced function definitions are still bound, dead code elimination removes them later
fun = itir.FunctionDefinition(id="fun", params=[im.sym("a")], expr=im.deref("a"))
testee = program_factory(
body=[itir.SetAt(expr=im.deref("inp"), target=im.ref("out"), domain=domain)],
function_definitions=[fun],
)
expected = program_factory(
body=[
itir.SetAt(
expr=im.let("fun", im.lambda_("a")(im.deref("a")))(im.deref("inp")),
target=im.ref("out"),
domain=domain,
)
],
function_definitions=[],
)

actual = inline_fundefs.inline_fundefs(testee)
assert actual == expected


def test_shadowed_by_binder():
# a binder of the same name shadows the function definition and must not be replaced by it
fun = itir.FunctionDefinition(id="fun", params=[im.sym("a")], expr=im.deref("a"))
stencil = im.lambda_("fun")(im.deref("fun"))
testee = program_factory(
body=[
itir.SetAt(
expr=im.as_fieldop(stencil, domain)("inp"), target=im.ref("out"), domain=domain
)
],
function_definitions=[fun],
)

expected = program_factory(
body=[
itir.SetAt(
expr=im.let("fun", im.lambda_("a")(im.deref("a")))(
# the `fun` binder of the stencil is untouched
im.as_fieldop(stencil, domain)("inp")
),
target=im.ref("out"),
domain=domain,
)
],
function_definitions=[],
)

actual = inline_fundefs.inline_fundefs(testee)
assert actual == expected


def test_dependent_fundefs():
# function definitions may reference each other, independent of their order in the program
fun1 = itir.FunctionDefinition(id="fun1", params=[im.sym("a")], expr=im.deref("a"))
fun2 = itir.FunctionDefinition(id="fun2", params=[im.sym("a")], expr=im.call("fun1")("a"))
testee = program_factory(
body=[itir.SetAt(expr=im.call("fun2")("inp"), target=im.ref("out"), domain=domain)],
function_definitions=[fun2, fun1],
)
expected = program_factory(
body=[
itir.SetAt(
expr=im.let("fun1", im.lambda_("a")(im.deref("a")))(
im.let("fun2", im.lambda_("a")(im.call("fun1")("a")))(im.call("fun2")("inp"))
),
target=im.ref("out"),
domain=domain,
)
],
function_definitions=[],
)

actual = inline_fundefs.inline_fundefs(testee)
assert actual == expected


def test_if_stmt():
fun = itir.FunctionDefinition(id="fun", params=[im.sym("a")], expr=im.deref("a"))
testee = program_factory(
body=[
itir.IfStmt(
cond=im.call("fun")(True),
true_branch=[
itir.SetAt(expr=im.call("fun")("inp"), target=im.ref("out"), domain=domain)
],
false_branch=[],
)
],
function_definitions=[fun],
)

actual = inline_fundefs.inline_fundefs(testee)
binding = im.let("fun", im.lambda_("a")(im.deref("a")))
assert actual.body[0].cond == binding(im.call("fun")(True))
assert actual.body[0].true_branch[0].expr == binding(im.call("fun")("inp"))
Loading