Skip to content
Closed
Show file tree
Hide file tree
Changes from 28 commits
Commits
Show all changes
91 commits
Select commit Hold shift + click to select a range
1b4707c
Tracer prototype
SF-N Feb 20, 2026
902f8a3
Merge branch 'main' into tracer_support
SF-N Apr 27, 2026
02f881f
Introduce GTIR tree_map builtin and transform to make_tuple, also sup…
SF-N Apr 27, 2026
0ec4692
Run pre-commit and fix some tests
SF-N Apr 27, 2026
ab84ecc
Run CollapseTuple after UnrollTreeMap
SF-N Apr 28, 2026
36d6956
Merge branch 'main' into tracer_support_tree_map
SF-N Apr 28, 2026
152300e
Address review comments
SF-N Apr 28, 2026
d459b0e
Address further review comments
SF-N Apr 28, 2026
5c4b018
Merge remote-tracking branch 'origin/main' into tracer_support
tehrengruber Apr 29, 2026
97af81e
Apply review comments
SF-N Apr 29, 2026
8d75708
Merge branch 'main' into tracer_support_tree_map
SF-N Apr 29, 2026
0dfc80c
Add support for nested tuples
tehrengruber Apr 30, 2026
067bc29
Merge branch 'main' into tracer_support_tree_map
SF-N May 4, 2026
b771d66
Improve testing, typing, formatting
tehrengruber May 5, 2026
08ed490
Fix name shadowing
tehrengruber May 6, 2026
9e23d2d
Cleanup
tehrengruber May 6, 2026
97235d9
More fixes and commentary
tehrengruber May 11, 2026
dd30833
Fix failing tests
tehrengruber May 11, 2026
15b233e
Add test for calling a fo from a tuple comprehension
tehrengruber May 11, 2026
4f0b5d4
Fix format
tehrengruber May 11, 2026
b27f80b
Fix format
tehrengruber May 11, 2026
b69d700
Merge remote-tracking branch 'origin/main' into tracer_support
tehrengruber May 11, 2026
24f4c90
Small fix
tehrengruber May 11, 2026
32e5b2d
Rename map_ -> map_list
SF-N May 28, 2026
a7175d7
Run pre-commit
SF-N May 28, 2026
4f89818
Merge branch 'main' into tracer_support_tree_map
SF-N May 28, 2026
2779fd0
Refactor tree_map_tuple and add map_tuple with unrolling support
SF-N May 28, 2026
80f3273
Rename
SF-N May 28, 2026
454e15f
Minor fix
SF-N May 28, 2026
55f1799
Minor fixes
SF-N May 28, 2026
12dfecb
Some more fixes wrt. Copilot review
SF-N May 28, 2026
fc6b1cb
Merge branch 'main' into tracer_support
SF-N May 28, 2026
8a1febd
Merge branch 'main' into tracer_support_tree_map
SF-N Jun 2, 2026
31b969a
Remove unnecessary CollapseTuple loop
SF-N Jun 2, 2026
c7fc102
Reposition UnrollTupleMaps and simplify CollapseTuple usage
SF-N Jun 2, 2026
7993b9c
Merge branch 'main' into tracer_support_tree_map
SF-N Jun 2, 2026
b7f8ba9
Refactor tree_map unrolling
SF-N Jun 16, 2026
3d38868
Cleanup
SF-N Jun 16, 2026
7d5c86c
Revert "Cleanup"
SF-N Jun 17, 2026
747f36e
Revert "Refactor tree_map unrolling"
SF-N Jun 17, 2026
b767700
Cleanup
SF-N Jun 17, 2026
e91f1f1
Merge branch 'origin-main' into tracer_support_tree_map
SF-N Jun 17, 2026
9f3474d
Enhance tuple comprehension handling with fixed-length mappers
SF-N Jun 18, 2026
7b270a3
Address review comment
SF-N Jun 18, 2026
d3d4e46
Remove CollapseTuple pass after UnrollTupleMaps
SF-N Jun 19, 2026
d0272df
Remove program wrapper in tests
SF-N Jun 19, 2026
56f234e
Merge branch 'tracer_support_tree_map' of github.com:SF-N/gt4py into …
SF-N Jun 19, 2026
b7bb0b2
Fix test
SF-N Jun 19, 2026
00e077c
Refactor TupleComprehension lowering
SF-N Jun 22, 2026
6a5a980
Fix mypy issues
SF-N Jun 22, 2026
e55253b
Cleanup
SF-N Jun 22, 2026
4fbac27
Further refactor to get rid of FixedTupleComprehension
SF-N Jun 23, 2026
7a41dda
Fix test_with_tuples_of_local_fields
SF-N Jun 23, 2026
e73cb27
Exclude fixed-length tuples with heterogeneous elements to avoid dupl…
SF-N Jun 23, 2026
dc354d9
Further refactoring
SF-N Jun 24, 2026
d864531
Remove not implemented case
SF-N Jun 24, 2026
08d06a1
Merge branch 'main' into tracer_support
SF-N Jun 24, 2026
158d540
Merge branch 'main' into tracer_support_tree_map
SF-N Jun 24, 2026
7d8f56c
Also allow itir.Expr in UnrollTupleMaps and run tye_inference when ne…
SF-N Jun 24, 2026
7808b0f
Merge branch 'tracer_support_tree_map' of github.com:SF-N/gt4py into …
SF-N Jun 24, 2026
556178f
Merge branch 'main' into tracer_support_tree_map
SF-N Jun 30, 2026
01e2754
Merge branch 'main' into tracer_support
SF-N Jun 30, 2026
aed9577
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 2, 2026
bbf0679
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 3, 2026
5102336
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 3, 2026
c6d5a2d
Address review comments
SF-N Jul 6, 2026
242343f
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 6, 2026
06d56ee
Merge branch 'tracer_support_tree_map' of github.com:SF-N/gt4py into …
SF-N Jul 6, 2026
3f24c06
Update test
SF-N Jul 6, 2026
904623b
Run pre-commit
SF-N Jul 6, 2026
a2da186
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 8, 2026
9205669
Remove special casing from UnrollTupleMaps and update tests
SF-N Jul 8, 2026
349d239
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 9, 2026
9b610b4
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 13, 2026
7e7bde5
Apply review comments
SF-N Jul 13, 2026
2c15948
Apply further review comments
SF-N Jul 13, 2026
fb724d3
Remove support for tre_map_tuple with multi-args
SF-N Jul 13, 2026
a7ff65a
Apply review comments
SF-N Jul 17, 2026
dc59b75
Rename filename unroll -> expand
SF-N Jul 17, 2026
3b0b4c3
Minor fixes
SF-N Jul 17, 2026
9ab7df4
Merge branch 'main' into tracer_support_tree_map
SF-N Jul 17, 2026
b96652e
Merge branch 'tracer_support_tree_map' into tracer_support
SF-N Jul 17, 2026
a670352
Fix merge issues
SF-N Jul 20, 2026
fced926
Merge branch 'main' into tracer_support
SF-N Jul 20, 2026
502ae5c
Merge branch 'main' into tracer_support
SF-N Jul 20, 2026
ef504e6
Merge branch 'main' into tracer_support
SF-N Jul 27, 2026
e9b8ec2
Merge commit 'bfe883c161cf229751cdc4e32713a00bdbbd7ac8' into tracer_s…
tehrengruber Aug 21, 2026
6d5aab9
Adapt branch code to conventions introduced on main
tehrengruber Aug 21, 2026
4957eb2
Merge branch 'main' of https://github.com/GridTools/gt4py into tracer…
tehrengruber Aug 25, 2026
d700c31
Remove 'target: Any' workaround in TupleComprehensionMapper
tehrengruber Aug 25, 2026
ccc0722
Merge branch 'main' into tracer_support
SF-N Aug 26, 2026
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
29 changes: 29 additions & 0 deletions src/gt4py/next/ffront/field_operator_ast.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,35 @@ class TupleExpr(Expr):
elts: list[Expr]


# TODO(tehrengruber): extend this to supported nested tuple comprehension.
# e.g. `tuple(element_expr for child in nested_tuple for grand_child in child)`
# would be represented by:
# ```
# class TupleComprehension(Expr): # ruff: noqa: ERA001
# inner: TupleComprehensionMapper | NestedTupleCompr # ruff: noqa: ERA001
# class NestedTupleCompr(Expr, SymbolTableTrait): # ruff: noqa: ERA001
# params: tuple[DataSymbol] # ruff: noqa: ERA001
# body: TupleComprehension # ruff: noqa: ERA001
# ```
class TupleComprehension(Expr):
"""
tuple(element_expr for target in iterable)
Note: The structure here differs from the one in the Python AST. Here we group target and
element expression in order to cleanly nest by the symbols being introduced, whereas in
the Python AST target and iterable are grouped into generator nodes.
"""

inner: TupleComprehensionMapper
iterable: Expr


# This is essentially a lambda. The difference is that for a lambda we might not know the type of
# the args; therefore this is named differently at the moment.
class TupleComprehensionMapper(LocatedNode, SymbolTableTrait):
target: Any # TODO(tehrengruber): should be NestedTuple[DataSymbol], but this breaks in eve
element_expr: Expr


class UnaryOp(Expr):
op: dialect_ast_enums.UnaryOperator
operand: Expr
Expand Down
95 changes: 93 additions & 2 deletions src/gt4py/next/ffront/foast_passes/type_deduction.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,13 @@
#
# Please, refer to the LICENSE file in the root directory.
# SPDX-License-Identifier: BSD-3-Clause

import textwrap
from typing import Any, Optional, Sequence, TypeAlias, TypeVar, cast

import gt4py.next.ffront.field_operator_ast as foast
from gt4py import eve
from gt4py.eve import NodeTranslator, NodeVisitor, traits
from gt4py.eve.extended_typing import NestedTuple
from gt4py.next import errors
from gt4py.next.common import Dimension, DimensionKind, promote_dims
from gt4py.next.ffront import (
Expand All @@ -24,6 +24,7 @@
from gt4py.next.ffront.foast_passes import utils as foast_utils
from gt4py.next.iterator import builtins
from gt4py.next.type_system import type_info, type_specifications as ts, type_translation
from gt4py.next.utils import tree_map


OperatorNodeT = TypeVar("OperatorNodeT", bound=foast.LocatedNode)
Expand Down Expand Up @@ -453,6 +454,10 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri
f"Tuples need to be indexed with literal integers, got '{node.index}'.",
) from ex
new_type = types[index]
case ts.VarArgType(element_type=element_type):
new_type = (
element_type # TODO: we only temporarily allow any index for vararg types

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is for direct access to tracers[0] * factor, tracers[1] * factor, which I personally think is an anti pattern. I left it here until we take a decision on this. We could also make it an optional feature. One of the disadvantages is that it is not possible to fully type check the field operator at definition time, since the tuple length is only known at call / compile time. The user will then get an error in unroll_map_tuple.

)
Comment on lines +462 to +464
Comment on lines +461 to +464
case ts.OffsetType(source=source, target=(target1, target2)):
if not target2.kind == DimensionKind.LOCAL:
raise errors.DSLError(
Expand Down Expand Up @@ -715,6 +720,90 @@ def visit_TupleExpr(self, node: foast.TupleExpr, **kwargs: Any) -> foast.TupleEx
new_type = ts.TupleType(types=[element.type for element in new_elts])
return foast.TupleExpr(elts=new_elts, type=new_type, location=node.location)

def visit_TupleComprehension(
self, node: foast.TupleComprehension, **kwargs: Any
) -> foast.TupleComprehension:
target = self.visit(node.inner.target, **kwargs)
iterable = self.visit(node.iterable, **kwargs)

def deduce_target_type(
target: NestedTuple[foast.Symbol] | foast.Symbol,
element_type: ts.TypeSpec,
inner_kwargs: dict[str, Any],
) -> NestedTuple[foast.Symbol] | foast.Symbol:
@tree_map(with_path_arg=True)
def process_target(target_el: foast.Symbol, path: tuple[int, ...]) -> foast.Symbol:
try:
type_ = element_type
for i in path:
if not isinstance(type_, ts.TupleType) or len(type_.types) <= i:
raise IndexError()
type_ = type_.types[i]
return self.visit(target_el, refine_type=type_, **inner_kwargs)
except IndexError:
raise errors.DSLError(
target_el.location, f"Cannot unpack non-iterable '{type_}' object."
) from None

return process_target(target)

def deduce_mapper(
element_type: ts.DataType,
) -> foast.TupleComprehensionMapper:
inner_kwargs = {**kwargs, "symtable": kwargs["symtable"].new_child()}
new_target = deduce_target_type(target, element_type, inner_kwargs)
return foast.TupleComprehensionMapper(
target=new_target,
element_expr=self.visit(node.inner.element_expr, **inner_kwargs),
location=node.location,
)

if isinstance(iterable.type, ts.TupleType):
if len(iterable.type.types) == 0:
raise errors.DSLError(
iterable.location,
"Cannot iterate over an empty tuple in a tuple comprehension.",
)
if not all(
isinstance(element_type, ts.DataType) for element_type in iterable.type.types
):
raise errors.DSLError(
iterable.location,
"Tuple comprehension iterable elements must be data types.",
)

element_types = cast(list[ts.DataType], iterable.type.types)
if not all(element_type == element_types[0] for element_type in element_types):
raise NotImplementedError(
"Tuple comprehensions over fixed-length tuples require all iterable "
"elements to have the same type."
)
new_mapper = deduce_mapper(element_types[0])
result = foast.TupleComprehension(
inner=new_mapper,
iterable=iterable,
location=node.location,
type=ts.TupleType(types=[new_mapper.element_expr.type for _ in element_types]),
)
return result
elif isinstance(iterable.type, ts.VarArgType):
element_type = iterable.type.element_type
Comment thread
SF-N marked this conversation as resolved.
new_mapper = deduce_mapper(element_type)
element_expr = new_mapper.element_expr
return_type = ts.VarArgType(element_type=element_expr.type)

return foast.TupleComprehension(
inner=new_mapper,
iterable=iterable,
location=node.location,
type=return_type,
)
else:
raise errors.DSLError(
iterable.location,
f"Iterable in generator expression must be a tuple, got '{iterable.type}'.",
)

def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call:
new_func = self.visit(node.func, **kwargs)
new_args = self.visit(node.args, **kwargs)
Expand Down Expand Up @@ -998,7 +1087,9 @@ def deduce_return_type(
f"Field arguments to '{func_name}' must be of same dtype, got '{t_dtype}' != "
f"'{f_dtype}'.",
)
return_dims = promote_dims(cond_dims, type_info.extract_dims(type_info.promote(tb, fb)))
return_dims = promote_dims(
cond_dims, type_info.extract_dims(tb), type_info.extract_dims(fb)
)
return_type = ts.FieldType(dims=return_dims, dtype=t_dtype)
return return_type

Expand Down
12 changes: 12 additions & 0 deletions src/gt4py/next/ffront/foast_pretty_printer.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,18 @@ def apply(cls, node: foast.LocatedNode, **kwargs: Any) -> str: # type: ignore[o

UnaryOp = as_fmt("{op}{operand}")

def visit_TupleComprehensionMapper(
self, node: foast.TupleComprehensionMapper, **kwargs: Any
) -> str:
element_expr = self.visit(node.element_expr, **kwargs)
target = self.visit(node.target, **kwargs)
return f"{element_expr} for {target}"

def visit_TupleComprehension(self, node: foast.TupleComprehension, **kwargs: Any) -> str:
mapper = self.visit(node.inner, **kwargs)
iterable = self.visit(node.iterable, **kwargs)
return f"tuple(({mapper} in {iterable}))"

def visit_UnaryOp(self, node: foast.UnaryOp, **kwargs: Any) -> str:
if node.op is dialect_ast_enums.UnaryOperator.NOT:
op = "not "
Expand Down
75 changes: 70 additions & 5 deletions src/gt4py/next/ffront/foast_to_gtir.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,12 @@


import dataclasses
import functools
import warnings
from typing import Any, Callable, Optional

from gt4py import eve
from gt4py.eve.extended_typing import Never, cast
from gt4py.eve.extended_typing import Never
from gt4py.next import utils
from gt4py.next.ffront import (
dialect_ast_enums,
Expand Down Expand Up @@ -260,6 +261,73 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> itir.Expr:
def visit_TupleExpr(self, node: foast.TupleExpr, **kwargs: Any) -> itir.Expr:
return im.make_tuple(*[self.visit(el, **kwargs) for el in node.elts])

def _bind_tuple_comprehension_target(
self,
comprehension_target: itir.Sym | tuple,
element_expr: itir.Expr,
iterable_element: itir.Expr | str,
) -> itir.Expr:
"""Return ``element_expr`` with the comprehension target bound to one element."""
# For `2.0 * local_el + scalar_el for local_el, scalar_el in iterable`:
# - `comprehension_target`: `(local_el, scalar_el)`
# - `element_expr`: `2.0 * local_el + scalar_el`
# - `iterable_element` is the current element from `iterable`
# Returns `let local_el = iterable_element[0], scalar_el = iterable_element[1]
# in element_expr`.
if not isinstance(comprehension_target, tuple):
return im.let(comprehension_target, iterable_element)(element_expr)

flat_targets = utils.flatten_nested_tuple(comprehension_target)
nested_target_values = utils.tree_map(
lambda _, path: functools.reduce(
lambda element, index: im.tuple_get(index, element), path, iterable_element
),
with_path_arg=True,
)(comprehension_target)

flat_target_values = utils.flatten_nested_tuple(nested_target_values) # type: ignore[arg-type]

target_bindings = tuple(zip(flat_targets, flat_target_values, strict=True))
return im.let(*target_bindings)(element_expr) # type: ignore[arg-type]

def visit_TupleComprehension(self, node: foast.TupleComprehension, **kwargs: Any) -> itir.Expr:
# e.g. tuple(2.0 * el for el in (a, a))` or `tuple(2.0 * el for el in (a(V2E), a(V2E)))`
# `tuple(2.0 * local_el + scalar_el for local_el, scalar_el in ((a(V2E), b), (c(V2E), d)))`.
# Only homogeneous (fixed-length and variable-length) tuples are supported.
comprehension_target = self.visit(node.inner.target, **kwargs)
element_expr = self.visit(node.inner.element_expr, **kwargs)
iterable_expr = self.visit(node.iterable, **kwargs)
iterable_type = node.iterable.type

def lower_body_for_iterable_element(iterable_element: itir.Expr | str) -> itir.Expr:
return self._bind_tuple_comprehension_target(
comprehension_target, element_expr, iterable_element
)

if isinstance(iterable_type, ts.TupleType):
assert isinstance(node.type, ts.TupleType)
iterable_value_name = next(self.uid_generator["__tuple_comprh"])

fixed_tuple_elements = [
lower_body_for_iterable_element(im.tuple_get(element_index, iterable_value_name))
for element_index in range(len(iterable_type.types))
]

result_tuple = im.make_tuple(*fixed_tuple_elements)
return im.let(iterable_value_name, iterable_expr)(result_tuple)

assert isinstance(iterable_type, ts.VarArgType)
assert isinstance(node.type, ts.VarArgType)
if not isinstance(comprehension_target, tuple):
map_tuple_lambda = im.lambda_(comprehension_target)(element_expr)
else:
iterable_element_param = next(self.uid_generator["__tuple_comprh"])
map_tuple_lambda = im.lambda_(iterable_element_param)(
lower_body_for_iterable_element(iterable_element_param)
)

return im.call(im.call("map_tuple")(map_tuple_lambda))(iterable_expr)

def visit_UnaryOp(self, node: foast.UnaryOp, **kwargs: Any) -> itir.Expr:
# TODO(tehrengruber): extend iterator ir to support unary operators
dtype = type_info.extract_dtype(node.type)
Expand Down Expand Up @@ -521,10 +589,7 @@ def _visit_type_constr(self, node: foast.Call, **kwargs: Any) -> itir.Expr:
return im.literal(str(val), target_type)

def _make_literal(self, val: Any, type_: ts.TypeSpec) -> itir.Expr:
if isinstance(type_, ts.COLLECTION_TYPE_SPECS):
type_ = cast(
ts.CollectionTypeSpec, type_
) # This shouldn't be needed after the previous isinstance() check
if isinstance(type_, (ts.TupleType, ts.NamedCollectionType)):
# This code-path is only active in the init of a scan,
# as otherwise the frontend generates tuple expressions of `Constant`s.
val = arguments.extract(val) if isinstance(type_, ts.NamedCollectionType) else val
Expand Down
3 changes: 2 additions & 1 deletion src/gt4py/next/ffront/foast_to_past.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,9 +113,10 @@ def __call__(self, inp: ConcreteFOASTOperatorDef) -> ConcretePASTProgramDef:
*partial_program_type.definition.kw_only_args.keys(),
]
assert isinstance(type_, ts.CallableType)
assert arg_types[-1] == type_info.return_type(
return_type = type_info.return_type(
type_, with_args=list(arg_types), with_kwargs=kwarg_types
)
assert type_info.is_concretizable(return_type, arg_types[-1])
assert args_names[-1] == "out"

params_decl: list[past.Symbol] = [
Expand Down
Loading