From 0cfc71684d6b5e4497160c9b327771fcec4bec1e Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Mon, 21 Sep 2026 08:35:42 +0000 Subject: [PATCH 1/2] fix[next]: lower function definitions to `let` bindings `InlineFundefs` replaced every `SymRef` matching a program-level function definition, regardless of any binder of the same name in between. A lambda parameter named like a function definition was therefore silently replaced by that function, producing an ill-typed program (e.g. an `AssertionError` in `type_synthesizer._canonicalize_nb_fields`). Instead of substituting references, bind all function definitions in a `let` wrapping every statement expression. Regular scoping rules then apply, i.e. an inner binder shadows the function definition, and unused bindings are removed by dead code elimination, which makes the separate `prune_unreferenced_fundefs` pass obsolete. --- .../iterator/transforms/inline_fundefs.py | 97 ++++++----- .../next/iterator/transforms/pass_manager.py | 7 +- .../ffront_tests/test_closure_vars.py | 19 +++ .../transforms_tests/test_inline_fundefs.py | 153 ++++++++++++++++++ 4 files changed, 233 insertions(+), 43 deletions(-) create mode 100644 tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_fundefs.py diff --git a/src/gt4py/next/iterator/transforms/inline_fundefs.py b/src/gt4py/next/iterator/transforms/inline_fundefs.py index 2b8767e4a2..8f9f675f10 100644 --- a/src/gt4py/next/iterator/transforms/inline_fundefs.py +++ b/src/gt4py/next/iterator/transforms/inline_fundefs.py @@ -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), ) diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index 067dde468f..fd628b3416 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -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. @@ -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( diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_closure_vars.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_closure_vars.py index 086cc4a282..2a527b5fa4 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_closure_vars.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_closure_vars.py @@ -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) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_fundefs.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_fundefs.py new file mode 100644 index 0000000000..c3fa309e3e --- /dev/null +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_fundefs.py @@ -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")) From 98ae1cbe55b1e5959f5271026534d4c7b0e8347b Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Mon, 21 Sep 2026 17:20:59 +0000 Subject: [PATCH 2/2] test[next]: make `flux` in `test_hdiff` a plain closure The traced `flux` function definition is polymorphic in its parameter `d`, it is called with an `I` and a `J` offset. Type inference can not type a `let` bound function used at several types, which the function definitions are lowered to now. Dropping the `fundef` decorator makes tracing inline `flux` as a lambda instead, which keeps the test unchanged otherwise. --- .../multi_feature_tests/iterator_tests/test_hdiff.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_hdiff.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_hdiff.py index 3ebacfd80e..b4956dbcf0 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_hdiff.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_hdiff.py @@ -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)