From e5b676acbb663b05bad412d8db90a4a46bc977f2 Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Thu, 27 Aug 2026 09:48:26 +0000 Subject: [PATCH 1/8] fix[next]: 'concat_where' crashes on local (neighbor-list) fields The GTIR 'concat_where' type synthesizer derives the result dims via 'type_info.promote(tb, fb)', which asserted that the promoted dtypes are 'ScalarType'. A local field carries a 'ListType' dtype at the GTIR level, so any 'concat_where' with a local-field branch failed with a bare 'AssertionError' during type inference -- and it did so even though both branches have the identical dtype, since the assertion is on the dtype's kind, not on the operands differing. Let 'promote' handle 'ListType' the same way it handles 'ScalarType': both promote only between equal types, so the two cases collapse into a single check. The docstring records that a 'ListType' only ever reaches 'promote' from the ITIR level, because the frontend represents the same concept as a field with a local dimension in 'dims' and a scalar dtype. The combination was previously untested ('test_concat_where.py' had no local-dimension coverage), so the latent assert never fired in CI. --- .../ffront/foast_passes/type_deduction.py | 8 ++- src/gt4py/next/type_system/type_info.py | 34 ++++++--- .../ffront_tests/test_concat_where.py | 72 ++++++++++++++++++- 3 files changed, 100 insertions(+), 14 deletions(-) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index c9f51ad080..6d25c3edb3 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -871,8 +871,12 @@ def _visit_math_built_in(self, node: foast.Call, **kwargs: Any) -> foast.Call: ) elif func_name in fbuiltins.BINARY_MATH_NUMBER_BUILTIN_NAMES: try: - return_type = type_info.promote( - *((cast(ts.FieldType | ts.ScalarType, arg.type)) for arg in node.args) + return_type = cast( + # a `ListType` only occurs at the ITIR level, never in the frontend + ts.FieldType | ts.ScalarType, + type_info.promote( + *((cast(ts.FieldType | ts.ScalarType, arg.type)) for arg in node.args) + ), ) except ValueError as ex: raise errors.DSLError(node.location, error_msg_preamble) from ex diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index ac8467a5f6..1dd404fde3 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -9,7 +9,7 @@ import functools import types from collections.abc import Callable, Collection, Iterable, Iterator -from typing import Any, Final, Literal, Sequence, Type, TypeGuard, TypeVar, cast, overload +from typing import Any, Final, Literal, Sequence, Type, TypeGuard, TypeVar, overload import numpy as np @@ -556,15 +556,20 @@ def is_concretizable(symbol_type: ts.TypeSpec, to_type: ts.TypeSpec) -> bool: def promote( - *types: ts.FieldType | ts.ScalarType, always_field: bool = False -) -> ts.FieldType | ts.ScalarType: + *types: ts.FieldType | ts.ScalarType | ts.ListType, always_field: bool = False +) -> ts.FieldType | ts.ScalarType | ts.ListType: """ - Promote a set of field or scalar types to a common type. + Promote a set of field, scalar or list types to a common type. The resulting type is defined on all dimensions of the arguments, respecting the individual order of the dimensions of each argument (see :func:`common.promote_dims` for more details). + A `ListType` is the dtype of a local (neighbor-list) field and only ever reaches this + function from the ITIR level. The frontend represents the same concept as a field with + a local dimension in `dims` and a scalar dtype, so it never passes a list here. Lists + promote exactly like scalars, i.e. only between equal types. + >>> dtype = ts.ScalarType(kind=ts.ScalarKind.INT64) >>> I, J, K = (common.Dimension(value=dim) for dim in ["I", "J", "K"]) >>> promoted: ts.FieldType = promote( @@ -578,18 +583,27 @@ def promote( ... ) >>> promoted.dims == [I, J, K] and promoted.dtype == dtype True + + >>> V2E = common.Dimension(value="V2E", kind=common.DimensionKind.LOCAL) + >>> list_dtype = ts.ListType(element_type=dtype, offset_type=V2E) + >>> promote( + ... ts.FieldType(dims=[I], dtype=list_dtype), + ... ts.FieldType(dims=[I, J], dtype=list_dtype), + ... ) == ts.FieldType(dims=[I, J], dtype=list_dtype) + True """ - if not always_field and all(isinstance(type_, ts.ScalarType) for type_ in types): + # Lists are only reached from the ITIR level (see above) and behave like scalars here: + # both promote only between equal types. + if not always_field and all(isinstance(type_, (ts.ScalarType, ts.ListType)) for type_ in types): if not all(type_ == types[0] for type_ in types): - raise ValueError("Could not promote scalars of different dtype (not implemented).") - if not all(type_.shape is None for type_ in types): # type: ignore[union-attr] + raise ValueError("Could not promote dtypes of different type (not implemented).") + if isinstance(types[0], ts.ScalarType) and types[0].shape is not None: raise NotImplementedError("Shape promotion not implemented.") return types[0] elif all(isinstance(type_, (ts.ScalarType, ts.FieldType)) for type_ in types): dims = common.promote_dims(*(extract_dims(type_) for type_ in types)) - extracted_dtypes = [extract_dtype(type_) for type_ in types] - assert all(isinstance(dtype, ts.ScalarType) for dtype in extracted_dtypes) - dtype = cast(ts.ScalarType, promote(*extracted_dtypes)) # type: ignore[arg-type] # checked is `ScalarType` + dtype = promote(*(extract_dtype(type_) for type_ in types)) + assert isinstance(dtype, (ts.ScalarType, ts.ListType)) return ts.FieldType(dims=dims, dtype=dtype) raise TypeError("Expected a 'FieldType' or 'ScalarType'.") diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py index 577f5a520e..f46c166e98 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py @@ -8,13 +8,23 @@ import numpy as np import pytest -from next_tests.integration_tests.cases import IDim, JDim, KDim, cartesian_case +from next_tests.integration_tests.cases import ( + E2V, + E2VDim, + Edge, + IDim, + JDim, + KDim, + cartesian_case, + unstructured_case, +) from gt4py import next as gtx -from gt4py.next import broadcast +from gt4py.next import broadcast, common, neighbor_sum from gt4py.next.ffront.experimental import concat_where from next_tests.integration_tests import cases from next_tests.integration_tests.cases_utils import ( exec_alloc_descriptor, + mesh_descriptor, ) pytestmark = pytest.mark.uses_concat_where @@ -390,6 +400,64 @@ def ref(interior0, boundary0, interior1, boundary1): cases.verify_with_default_data(cartesian_case, testee, ref) +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_sparse_fields +def test_with_local_field(unstructured_case, static_domains: bool): + @gtx.field_operator(static_domains=static_domains) + def testee(a: cases.VField, b: cases.VField) -> cases.EField: + t = concat_where(Edge < 2, a(E2V), b(E2V)) + return neighbor_sum(t, axis=E2VDim) + + e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + edge_mask = np.arange(unstructured_case.default_sizes[Edge]) < 2 + cases.verify_with_default_data( + unstructured_case, + testee, + ref=lambda a, b: np.sum( + np.where(edge_mask[:, np.newaxis], a[e2v_table], b[e2v_table]), + axis=1, + initial=0, + where=e2v_table != common._DEFAULT_SKIP_VALUE, + ), + ) + + +@pytest.mark.uses_tuple_returns +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_sparse_fields +def test_with_tuples_of_local_fields(unstructured_case, static_domains: bool): + @gtx.field_operator(static_domains=static_domains) + def testee( + a: cases.VField, + b: cases.VField, + c: cases.VField, + d: cases.VField, + ) -> tuple[cases.EField, cases.EField]: + t = concat_where(Edge < 2, (a(E2V), b(E2V)), (c(E2V), d(E2V))) + return neighbor_sum(t[0], axis=E2VDim), neighbor_sum(t[1], axis=E2VDim) + + e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + edge_mask = np.arange(unstructured_case.default_sizes[Edge]) < 2 + cases.verify_with_default_data( + unstructured_case, + testee, + ref=lambda a, b, c, d: ( + np.sum( + np.where(edge_mask[:, np.newaxis], a[e2v_table], c[e2v_table]), + axis=1, + initial=0, + where=e2v_table != common._DEFAULT_SKIP_VALUE, + ), + np.sum( + np.where(edge_mask[:, np.newaxis], b[e2v_table], d[e2v_table]), + axis=1, + initial=0, + where=e2v_table != common._DEFAULT_SKIP_VALUE, + ), + ), + ) + + def test_nested_conditions_with_empty_branches(cartesian_case, static_domains: bool): @gtx.field_operator(static_domains=static_domains) def testee(interior: cases.IField, boundary: cases.IField, N: gtx.int32) -> cases.IField: From 8643e6d67e8d3b63198bf7a0a2c082724be99d90 Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Thu, 27 Aug 2026 19:58:03 +0000 Subject: [PATCH 2/8] Add 'concat_where' test with nested tuple branches (moved from #2833) --- .../ffront_tests/test_concat_where.py | 45 +++++++++++++++++++ 1 file changed, 45 insertions(+) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py index f46c166e98..6ed7b244c4 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py @@ -400,6 +400,51 @@ def ref(interior0, boundary0, interior1, boundary1): cases.verify_with_default_data(cartesian_case, testee, ref) +@pytest.mark.uses_tuple_returns +def test_with_nested_tuples(cartesian_case, static_domains: bool): + @gtx.field_operator(static_domains=static_domains) + def testee( + interior0: cases.IJKField, + boundary0: cases.IJField, + interior1: cases.IJKField, + boundary1: cases.IJField, + interior2: cases.IJKField, + boundary2: cases.IJField, + ) -> tuple[cases.IJKField, tuple[cases.IJKField, cases.IJKField]]: + return concat_where( + KDim == 0, + (boundary0, (boundary1, boundary2)), + (interior0, (interior1, interior2)), + ) + + interiors = tuple(cases.allocate(cartesian_case, testee, f"interior{i}")() for i in range(3)) + boundaries = tuple(cases.allocate(cartesian_case, testee, f"boundary{i}")() for i in range(3)) + out = cases.allocate(cartesian_case, testee, cases.RETURN)() + + k = np.arange(0, cartesian_case.default_sizes[KDim]) + refs = tuple( + np.where( + k[np.newaxis, np.newaxis, :] == 0, + boundary.asnumpy()[:, :, np.newaxis], + interior.asnumpy(), + ) + for boundary, interior in zip(boundaries, interiors) + ) + + cases.verify( + cartesian_case, + testee, + interiors[0], + boundaries[0], + interiors[1], + boundaries[1], + interiors[2], + boundaries[2], + out=out, + ref=(refs[0], (refs[1], refs[2])), + ) + + @pytest.mark.uses_unstructured_shift @pytest.mark.uses_sparse_fields def test_with_local_field(unstructured_case, static_domains: bool): From 7b21e9c870660b8035870b5e8770cc71bda3f407 Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Thu, 27 Aug 2026 20:09:16 +0000 Subject: [PATCH 3/8] feat[next]: Add support for tuple comprehensions --- .../0028-Homogeneous_Tuple_Comprehensions.md | 41 ++++ docs/development/ADRs/next/README.md | 1 + src/gt4py/next/ffront/dialect_parser.py | 5 +- src/gt4py/next/ffront/field_operator_ast.py | 30 +++ .../ffront/foast_passes/type_deduction.py | 113 +++++++++++ src/gt4py/next/ffront/foast_pretty_printer.py | 12 ++ src/gt4py/next/ffront/foast_to_gtir.py | 70 +++++++ src/gt4py/next/ffront/foast_to_past.py | 5 +- src/gt4py/next/ffront/func_to_foast.py | 80 ++++++-- src/gt4py/next/ffront/lowering_utils.py | 22 ++- .../next/ffront/past_passes/type_deduction.py | 4 +- src/gt4py/next/type_system/type_info.py | 8 + .../next/type_system/type_specifications.py | 9 + .../next/type_system/type_translation.py | 53 ++++-- tests/next_tests/integration_tests/cases.py | 11 ++ .../feature_tests/ffront_tests/test_tuples.py | 99 +++++++++- .../feature_tests/ffront_tests/test_where.py | 70 ++++++- .../ffront_tests/test_diagnostic_messages.py | 35 ++++ .../ffront_tests/test_foast_to_gtir.py | 179 +++++++++++++++++- .../ffront_tests/test_func_to_foast.py | 33 ++++ .../test_type_translation.py | 20 +- 21 files changed, 840 insertions(+), 60 deletions(-) create mode 100644 docs/development/ADRs/next/0028-Homogeneous_Tuple_Comprehensions.md diff --git a/docs/development/ADRs/next/0028-Homogeneous_Tuple_Comprehensions.md b/docs/development/ADRs/next/0028-Homogeneous_Tuple_Comprehensions.md new file mode 100644 index 0000000000..9b3b34e4f3 --- /dev/null +++ b/docs/development/ADRs/next/0028-Homogeneous_Tuple_Comprehensions.md @@ -0,0 +1,41 @@ +--- +tags: [] +--- + +# Homogeneous Tuple Comprehensions + +- **Status**: valid +- **Authors**: Till Ehrengruber (@tehrengruber), Sara Faghih-Naini (@SF-N) +- **Created**: 2026-08-27 +- **Updated**: 2026-08-27 + +In the context of tuple comprehensions in the field-view frontend, facing the constraint that every FOAST node carries exactly one type, we decided to support only homogeneous iterables — all elements of the same type — and reject heterogeneous ones in the type deduction. + +## Context + +Tuple comprehensions, e.g. `tuple(2.0 * el for el in (a, b))`, are typed and lowered with a single mapper: one target (`el`) and one element expression (`2.0 * el`), shared by all elements of the iterable. In FOAST every node has a single `type` attribute, so the target symbol — and consequently every node in the element expression — can only be typed once. If the iterable's elements had different types, the mapper would need a different type per element, i.e. per-element re-typing (monomorphization) of the element expression, which the FOAST type system does not support. + +The same constraint exists at the GTIR level: the ITIR type inference also stores a single type per node (and asserts on conflicting re-assignment), so the single `map_tuple` lambda used to lower comprehensions over variable-length tuples can only have one function type. For variable-length iterables heterogeneity cannot occur in the first place, since `VarArgType` describes all elements with a single element type. Rejecting heterogeneous iterables in the FOAST type deduction just surfaces the error earliest, with a source location. + +## Decision + +Only homogeneous iterables are supported, both fixed-length and variable-length. Heterogeneous ones are rejected in the type deduction. E.g., with `a`, `b`, `c`, `d` of equal type: + +```python +tuple(2.0 * el for el in (a, b)) # supported +tuple(2.0 * el for el in (a(V2E), b(V2E))) # supported +tuple(local_el + el for local_el, el in ((a(V2E), b), (c(V2E), d))) # supported +tuple(2.0 * el for el in (a(V2E), b)) # rejected: local vs. non-local element +``` + +Note that homogeneity applies to the iterable's elements as a whole: in the third example each element is a pair of a local and a non-local field, but all elements share that same tuple type, so each target symbol still has a single consistent type. + +## Consequences + +- Typing and lowering stay simple: the element expression is visited once, with one type per node. +- Computations over differently-typed elements cannot be written as a comprehension; they must be spelled out per element. +- The restriction could be lifted for fixed-length iterables by typing the mapper generically: the target symbol gets a type variable bounded by the valid element types, so a single type per node still suffices during type deduction. After `map_tuple` expansion the mapper is instantiated once per element, and each instance can then be specialized to its concrete element type. This is possible as a follow-up without breaking existing code, since it only widens the set of accepted programs. For variable-length iterables there is nothing to lift: heterogeneity cannot occur, as `VarArgType` has a single element type by construction. + +## References + +- PR [#2833](https://github.com/GridTools/gt4py/pull/2833) (tuple comprehension support) diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index 1bf9d21812..45b6fd9240 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -29,6 +29,7 @@ Writing a new ADR is simple: - [0002 - Field View Lowering](0002-Field_View_Lowering.md) - [0010 - Domain in Field View](0010-Domain_in_Field_View.md) - [0013 - Scalar vs 0d-Fields](0013-Scalar_vs_0d_Fields.md) +- [0028 - Homogeneous Tuple Comprehensions](0028-Homogeneous_Tuple_Comprehensions.md) ### Iterator IR #iterator diff --git a/src/gt4py/next/ffront/dialect_parser.py b/src/gt4py/next/ffront/dialect_parser.py index a7c27a3181..c464ffd4dc 100644 --- a/src/gt4py/next/ffront/dialect_parser.py +++ b/src/gt4py/next/ffront/dialect_parser.py @@ -55,7 +55,10 @@ ), ast.SetComp: ("set comprehension", ()), ast.DictComp: ("dictionary comprehension", ()), - ast.GeneratorExp: ("generator expression", ()), + ast.GeneratorExp: ( + "generator expression", + ("Generator expressions are only supported as the argument of 'tuple(...)'.",), + ), ast.Lambda: ( "'lambda' expression", ("Define a separate function decorated with '@field_operator' instead.",), diff --git a/src/gt4py/next/ffront/field_operator_ast.py b/src/gt4py/next/ffront/field_operator_ast.py index f84dd53d62..f95f9f285a 100644 --- a/src/gt4py/next/ffront/field_operator_ast.py +++ b/src/gt4py/next/ffront/field_operator_ast.py @@ -12,6 +12,7 @@ from gt4py import eve from gt4py.eve import Coerced, Node, SourceLocation, SymbolName, SymbolRef, datamodels +from gt4py.eve.extended_typing import MaybeNestedInTuple from gt4py.eve.traits import SymbolTableTrait from gt4py.eve.type_definitions import StrEnum from gt4py.next import utils @@ -98,6 +99,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: MaybeNestedInTuple[DataSymbol] + element_expr: Expr + + class UnaryOp(Expr): op: dialect_ast_enums.UnaryOperator operand: Expr diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 6d25c3edb3..9e52486921 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -12,6 +12,7 @@ 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 MaybeNestedInTuple from gt4py.next import common, errors from gt4py.next.common import Dimension, DimensionKind, promote_dims from gt4py.next.ffront import ( @@ -24,6 +25,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) @@ -743,6 +745,117 @@ 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 _deduce_tuple_comprehension_target_type( + self, + target: MaybeNestedInTuple[foast.Symbol], + element_type: ts.TypeSpec, + **kwargs: Any, + ) -> MaybeNestedInTuple[foast.Symbol]: + """ + Deduce the types of the loop target, e.g. `(a, b)` in `tuple(a + b for (a, b) in it)`. + + Each target symbol starts out with a deferred type; revisiting it with `refine_type` + replaces that by the constituent of `element_type` at the symbol's unpacking position. + """ + + @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_, **kwargs) + except IndexError: + raise errors.DSLError( + target_el.location, f"Cannot unpack non-iterable '{type_}' object." + ) from None + + return process_target(target) + + def _deduce_tuple_comprehension_mapper( + self, + node: foast.TupleComprehension, + target: MaybeNestedInTuple[foast.Symbol], + element_type: ts.DataType, + **kwargs: Any, + ) -> foast.TupleComprehensionMapper: + """ + Deduce the per-element part, e.g. `a + b for (a, b)` in `tuple(a + b for (a, b) in it)`. + + Refines the target symbols against `element_type` and then types the element + expression with those symbols in scope (a fresh child symbol table). + """ + inner_kwargs = {**kwargs, "symtable": kwargs["symtable"].new_child()} + new_target = self._deduce_tuple_comprehension_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, + ) + + def visit_TupleComprehension( + self, node: foast.TupleComprehension, **kwargs: Any + ) -> foast.TupleComprehension: + # The target symbols are visited in `_deduce_tuple_comprehension_target_type`, + # where their types are refined against the iterable's element type. + target = node.inner.target + iterable = self.visit(node.iterable, **kwargs) + + 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) + # Only homogeneous iterables are supported, see ADR 0028. + 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 (see ADR 0028)." + ) + new_mapper = self._deduce_tuple_comprehension_mapper( + node, target, element_types[0], **kwargs + ) + 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 + new_mapper = self._deduce_tuple_comprehension_mapper( + node, target, element_type, **kwargs + ) + 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) diff --git a/src/gt4py/next/ffront/foast_pretty_printer.py b/src/gt4py/next/ffront/foast_pretty_printer.py index 8b2e369501..79826983a3 100644 --- a/src/gt4py/next/ffront/foast_pretty_printer.py +++ b/src/gt4py/next/ffront/foast_pretty_printer.py @@ -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 " diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 480a05812f..283253bce7 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -8,6 +8,7 @@ import dataclasses +import functools from typing import Any, Callable, Optional from gt4py import eve @@ -259,6 +260,75 @@ 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: + """ + Wrap `element_expr` in a `let` binding the `comprehension_target` to `iterable_element`. + + For `2.0 * a + b for (a, b) in iterable`: + - `comprehension_target`: `(a, b)` + - `element_expr`: `2.0 * a + b` + - `iterable_element`: the current element of `iterable` + + returns + `let a = iterable_element[0], b = iterable_element[1] in element_expr`. + """ + if isinstance(comprehension_target, itir.Sym): + 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: + # Only homogeneous iterables — all elements of the same type — are supported; + # heterogeneous ones are rejected in the type deduction (see ADR 0028). + 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 + + lower_body_for_iterable_element = functools.partial( + self._bind_tuple_comprehension_target, comprehension_target, element_expr + ) + + 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 isinstance(comprehension_target, itir.Sym): + 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) diff --git a/src/gt4py/next/ffront/foast_to_past.py b/src/gt4py/next/ffront/foast_to_past.py index 9a560b7ff8..1bbd3d9991 100644 --- a/src/gt4py/next/ffront/foast_to_past.py +++ b/src/gt4py/next/ffront/foast_to_past.py @@ -113,9 +113,12 @@ 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 ) + # Not equality: variadic comprehensions give a `VarArg[...]` return type, + # while 'out' is a concrete tuple. + assert type_info.is_concretizable(return_type, arg_types[-1]) assert args_names[-1] == "out" params_decl: list[past.Symbol] = [ diff --git a/src/gt4py/next/ffront/func_to_foast.py b/src/gt4py/next/ffront/func_to_foast.py index 14dceb25d1..d2926a90fc 100644 --- a/src/gt4py/next/ffront/func_to_foast.py +++ b/src/gt4py/next/ffront/func_to_foast.py @@ -11,9 +11,9 @@ import ast import textwrap import typing -from typing import Any, Type import gt4py.eve as eve +from gt4py.eve.extended_typing import Any, MaybeNestedInTuple from gt4py.next import errors from gt4py.next.ffront import ( dialect_ast_enums, @@ -324,7 +324,7 @@ def visit_Assign( if not isinstance(target, ast.Name): raise errors.DSLError(self.get_location(node), "Can only assign to names.") new_value = self.visit(node.value) - constraint_type: Type[ts.DataType] = ts.DataType + constraint_type: type[ts.DataType] = ts.DataType if isinstance(new_value, foast.TupleExpr): constraint_type = ts.TupleType elif ( @@ -542,24 +542,76 @@ def visit_NotEq(self, node: ast.NotEq, **kwargs: Any) -> foast.CompareOperator: return foast.CompareOperator.NOTEQ def _verify_builtin_type_constructor(self, node: ast.Call) -> None: - if len(node.args) > 0: - arg = node.args[0] - if not ( - isinstance(arg, ast.Constant) - or (isinstance(arg, ast.UnaryOp) and isinstance(arg.operand, ast.Constant)) - ): - raise errors.DSLError( - self.get_location(node), - f"'{self._func_name(node)}()' only takes literal arguments.", - ) + assert isinstance(node.func, ast.Name) + (arg,) = node.args + if not ( + isinstance(arg, ast.Constant) + or (isinstance(arg, ast.UnaryOp) and isinstance(arg.operand, ast.Constant)) + or (node.func.id == "tuple" and isinstance(arg, ast.GeneratorExp)) + ): + allowed = ( + "literal arguments or a generator expression" + if node.func.id == "tuple" + else "literal arguments" + ) + raise errors.DSLError( + self.get_location(node), + f"'{self._func_name(node)}()' only takes {allowed}.", + ) def _func_name(self, node: ast.Call) -> str: return node.func.id # type: ignore[attr-defined] # We want this to fail if the attribute does not exist unexpectedly. - def visit_Call(self, node: ast.Call, **kwargs: Any) -> foast.Call: - # TODO(tehrengruber): is this still needed or redundant with the checks in type deduction? + def visit_Call(self, node: ast.Call, **kwargs: Any) -> foast.Call | foast.TupleComprehension: if isinstance(node.func, ast.Name): func_name = self._func_name(node) + + if ( + func_name == "tuple" + and len(node.args) == 1 + and isinstance(gen_expr := node.args[0], ast.GeneratorExp) + ): + if len(gen_expr.generators) != 1: + raise errors.DSLError( + self.get_location(gen_expr), + "Nested generator expressions are not supported.", + hints=["Use a single 'for' clause iterating over one tuple."], + ) + if gen_expr.generators[0].ifs != []: + raise errors.DSLError( + self.get_location(gen_expr.generators[0].ifs[0]), + "Conditionals are not supported in generator expressions.", + notes=[ + ( + "The length of the resulting tuple must be known at compile time, " + "but an 'if' filter makes it depend on runtime values." + ) + ], + ) + + def parse_target(target: ast.expr) -> MaybeNestedInTuple[foast.DataSymbol]: + if isinstance(target, ast.Tuple): + return tuple(parse_target(el) for el in target.elts) + assert isinstance(target, ast.Name) and isinstance(target.ctx, ast.Store) + return foast.DataSymbol( + id=target.id, + location=self.get_location(target), + type=ts.DeferredType(constraint=None), + ) + + target = parse_target(gen_expr.generators[0].target) + + return foast.TupleComprehension( + inner=foast.TupleComprehensionMapper( + target=target, + element_expr=self.visit(gen_expr.elt, **kwargs), + location=self.get_location(node), + ), + iterable=self.visit(gen_expr.generators[0].iter, **kwargs), + location=self.get_location(node), + ) + + # TODO(tehrengruber): is this still needed or redundant with the checks in type deduction? if func_name in fbuiltins.TYPE_BUILTIN_NAMES: self._verify_builtin_type_constructor(node) diff --git a/src/gt4py/next/ffront/lowering_utils.py b/src/gt4py/next/ffront/lowering_utils.py index c4e35c9e18..f7a5d26991 100644 --- a/src/gt4py/next/ffront/lowering_utils.py +++ b/src/gt4py/next/ffront/lowering_utils.py @@ -20,6 +20,7 @@ def process_elements( objs: itir.Expr | Iterable[itir.Expr], current_el_type: ts.TypeSpec, arg_types: Optional[Iterable[ts.TypeSpec]] = None, + with_path_arg: bool = False, ) -> itir.FunCall: """ Recursively applies a processing function to all primitive constituents of a tuple or @@ -34,6 +35,7 @@ def process_elements( arg_types: If provided, a tuple of the type of each argument is passed to `process_func` as last argument. Note, that `arg_types` might coincide with `(current_el_type,)*len(objs)`, but not necessarily, in case of implicit broadcasts. + with_path_arg: If true, the index path to the current leaf is passed as last argument. """ if isinstance(objs, itir.Expr): objs = (objs,) @@ -47,6 +49,8 @@ def process_elements( tuple(im.ref(var_name) for var_name in var_names), current_el_type, arg_types=arg_types, + path=(), + with_path_arg=with_path_arg, ) return im.let(*bound_vars.items())(body) @@ -60,6 +64,8 @@ def _process_elements_impl( _current_el_exprs: Iterable[T], current_el_type: ts.TypeSpec, arg_types: Optional[Iterable[ts.TypeSpec]], + path: tuple[int, ...], + with_path_arg: bool, ) -> itir.Expr: if isinstance(current_el_type, (ts.TupleType, ts.NamedCollectionType)): result = im.make_tuple( @@ -70,17 +76,23 @@ def _process_elements_impl( im.tuple_get(i, current_el_expr) for current_el_expr in _current_el_exprs ), current_el_type.types[i], - arg_types=tuple(arg_t.types[i] for arg_t in arg_types) # type: ignore[attr-defined] # guaranteed by the requirement that `current_el_type` and each element of `arg_types` have the same tuple structure - if arg_types is not None - else None, + arg_types=( + tuple(arg_t.types[i] for arg_t in arg_types) # type: ignore[attr-defined] # guaranteed by the requirement that `current_el_type` and each element of `arg_types` have the same tuple structure + if arg_types is not None + else None + ), + path=(*path, i), + with_path_arg=with_path_arg, ) for i in range(len(current_el_type.types)) ) ) else: if arg_types is not None: - result = process_func(*_current_el_exprs, arg_types) + result = process_func( + *_current_el_exprs, arg_types, *((path,) if with_path_arg else ()) + ) else: - result = process_func(*_current_el_exprs) + result = process_func(*_current_el_exprs, *((path,) if with_path_arg else ())) return result diff --git a/src/gt4py/next/ffront/past_passes/type_deduction.py b/src/gt4py/next/ffront/past_passes/type_deduction.py index 9d021ceb51..fcc75b7374 100644 --- a/src/gt4py/next/ffront/past_passes/type_deduction.py +++ b/src/gt4py/next/ffront/past_passes/type_deduction.py @@ -248,7 +248,9 @@ def visit_Call(self, node: past.Call, **kwargs: Any) -> past.Call: operator_return_type = type_info.return_type( new_func.type, with_args=arg_types, with_kwargs=kwarg_types ) - if operator_return_type != new_kwargs["out"].type: + # Not equality: variadic comprehensions give a `VarArg[...]` return type, + # while 'out' is a concrete tuple. + if not type_info.is_compatible_type(operator_return_type, new_kwargs["out"].type): raise ValueError( "Expected keyword argument 'out' to be of " f"type '{operator_return_type}', got " diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 1dd404fde3..e3651b1c62 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -550,6 +550,14 @@ def is_concretizable(symbol_type: ts.TypeSpec, to_type: ts.TypeSpec) -> bool: or issubclass(type_class(to_type), symbol_type.constraint) ): return True + if isinstance(symbol_type, ts.VarArgType) and isinstance(to_type, ts.VarArgType): + return is_concretizable(symbol_type.element_type, to_type.element_type) + if isinstance(symbol_type, ts.VarArgType) and isinstance(to_type, ts.TupleType): + if len(to_type.types) == 0 or ( + all(type_ == to_type.types[0] for type_ in to_type.types) + and is_concretizable(symbol_type.element_type, to_type.types[0]) + ): + return True elif is_concrete(symbol_type): return symbol_type == to_type return False diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 59ac40f0f3..1016eba34e 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -148,6 +148,15 @@ def __len__(self) -> int: return len(self.types) +class VarArgType(DataType): + """Represents a variable number of arguments of the same type.""" + + element_type: DataType + + def __str__(self) -> str: + return f"VarArg[{self.element_type}]" + + class AnyPythonType: """Marker type representing any Python type which cannot be used for instantiation. diff --git a/src/gt4py/next/type_system/type_translation.py b/src/gt4py/next/type_system/type_translation.py index 850013ddfa..4ff4d3970b 100644 --- a/src/gt4py/next/type_system/type_translation.py +++ b/src/gt4py/next/type_system/type_translation.py @@ -75,6 +75,16 @@ def make_constructor_type(type_spec: ts.TypeSpec) -> ts.ConstructorType: ) ) + case ts.DeferredType(constraint=ts.TupleType): + return ts.ConstructorType( + definition=ts.FunctionType( + pos_only_args=[ts.DeferredType(constraint=None)], + pos_or_kw_args={}, + kw_only_args={}, + returns=ts.DeferredType(constraint=ts.VarArgType), + ) + ) + case ts.NamedCollectionType() as named_tuple_type: type_ = pkgutil.resolve_name(named_tuple_type.original_python_type) pos_or_kw_args = {k: t for k, t in zip(type_spec.keys, type_spec.types)} @@ -147,7 +157,7 @@ def canonicalize_type_hint( *, globalns: Optional[dict[str, Any]] = None, localns: Optional[dict[str, Any]] = None, -) -> tuple[Any, tuple[Any, ...], tuple[Any, ...]]: +) -> tuple[Any, tuple[Any, ...] | None, tuple[Any, ...]]: """ Canonicalize python type annotations as a tuple of (canonical_type, type_args, annotated_extra_args). """ @@ -178,7 +188,8 @@ def canonicalize_type_hint( type_hint = _resolve_type_alias(type_hint) canonical_type = typing.get_origin(type_hint) or type_hint - args = typing.get_args(type_hint) + # In order to distinguish `tuple` from `tuple[()]`, the former returns None here. + args = typing.get_args(type_hint) if typing.get_origin(type_hint) else None return canonical_type, args, tuple(extra_args) @@ -199,17 +210,37 @@ def from_type_hint( match canonical_type: case builtins.tuple: - if not args: - raise ValueError(f"Tuple annotation '{type_hint}' requires at least one argument.") - if Ellipsis in args: - raise ValueError(f"Unbound tuples '{type_hint}' are not allowed.") - tuple_types = [from_type_hint_same_ns(arg) for arg in args] - assert all(isinstance(elem, ts.DataType) for elem in tuple_types) - return ts.TupleType(types=tuple_types) + # Fixed-length tuple, e.g. `tuple[int32, float64]` + if ( + isinstance(args, tuple) + and len(args) > 0 + and not any(arg is Ellipsis for arg in args) + ): + tuple_types = [from_type_hint_same_ns(arg) for arg in args] + assert all(isinstance(elem, ts.DataType) for elem in tuple_types) + return ts.TupleType(types=tuple_types) + # Variable-length tuple, e.g. `tuple[int32, ...]` + elif isinstance(args, tuple) and len(args) == 2 and args[1] is Ellipsis: + return ts.VarArgType(element_type=from_type_hint_same_ns(args[0])) + # Unparametrized annotation, i.e. `tuple` or `tuple[()]` + elif args is None or (isinstance(args, tuple) and len(args) == 0): + # TODO(tehrengruber): We use `DeferredType` until we have an actual representation + # for a generic type. + return ts.DeferredType(constraint=ts.TupleType) + else: + raise ValueError( + f"Tuple annotation '{type_hint}' must either " + f"be a list of concrete arguments (e.g. 'tuple[int]'), " + f"be a variadic tuple (e.g. 'tuple[int, ...]'), " + f"or have no arguments (e.g. 'tuple')." + ) case common.Field: - if (n_args := len(args)) != 2: - raise ValueError(f"Field type requires two arguments, got {n_args}: '{type_hint}'.") + # `args` is `None` for an unparametrized `Field` annotation + if args is None or len(args) != 2: + raise ValueError( + f"Field type requires two arguments, got {len(args or ())}: '{type_hint}'." + ) dims: list[common.Dimension] = [] dim_arg, dtype_arg = args dim_arg = ( diff --git a/tests/next_tests/integration_tests/cases.py b/tests/next_tests/integration_tests/cases.py index 08fb817856..508125c341 100644 --- a/tests/next_tests/integration_tests/cases.py +++ b/tests/next_tests/integration_tests/cases.py @@ -616,6 +616,15 @@ def _allocate_from_type( for t in types ) ) + case ts.VarArgType(element_type=element_type): + return tuple( + ( + _allocate_from_type( + case=case, arg_type=t, domain=domain, dtype=dtype, strategy=strategy + ) + for t in [element_type] * 3 # TODO: revisit + ) + ) case ts.NamedCollectionType(types=types) as named_collection_type_spec: container_constructor = ( named_collections.make_named_collection_constructor_from_type_spec( @@ -661,6 +670,8 @@ def get_param_size(param_type: ts.TypeSpec, sizes: dict[gtx.Dimension, int]) -> return sum([get_param_size(t, sizes=sizes) for t in types]) case ts.NamedCollectionType(types=types): return sum([get_param_size(t, sizes=sizes) for t in types]) + case ts.VarArgType(element_type=element_type): + return get_param_size(ts.TupleType(types=[element_type] * 3), sizes) # TODO: revisit case _: raise TypeError(f"Can not get size for parameter of type '{param_type}'.") diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index 88a5a18af1..090f6a7e7e 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -10,14 +10,7 @@ import pytest import gt4py.next as gtx -from gt4py.next import ( - broadcast, - errors, - float64, - int32, - neighbor_sum, - utils as gt_utils, -) +from gt4py.next import broadcast, errors, float64, int32, neighbor_sum, utils as gt_utils from next_tests.integration_tests import cases from next_tests.integration_tests.cases import ( @@ -120,6 +113,94 @@ def testee(a: tuple[cases.IField, cases.IJField]) -> cases.IJField: ) +@pytest.mark.uses_tuple_args +def test_fixed_len_tuple_comprehension(cartesian_case): + @gtx.field_operator + def testee( + tracers: tuple[cases.IField, cases.IField], factor: int32 + ) -> tuple[cases.IField, cases.IField]: + return tuple(tracer * factor for tracer in tracers) + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda t, f: tuple(el * f for el in t), + ) + + +@pytest.mark.uses_tuple_args +def test_var_len_tuple_comprehension(cartesian_case): + @gtx.field_operator + def testee(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IField, ...]: + return tuple(tracer * factor for tracer in tracers) + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda t, f: tuple(el * f for el in t), + ) + + +@pytest.mark.uses_tuple_args +def test_tuple_comprehension_other_fo(cartesian_case): + @gtx.field_operator + def inner(tracer: cases.IField, factor: int32) -> cases.IField: + return tracer * factor + + @gtx.field_operator + def testee(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IField, ...]: + return tuple(inner(tracer, factor) for tracer in tracers) + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda t, f: tuple(el * f for el in t), + ) + + +@pytest.mark.uses_tuple_args +def test_nested_tuple_comprehension(cartesian_case): + @gtx.field_operator + def testee( + vals: tuple[tuple[cases.IField, ...], ...], factor: int32 + ) -> tuple[tuple[cases.IField, ...], ...]: + return tuple(tuple(grand_child * factor for grand_child in child) for child in vals) + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda t, f: tuple(tuple(grand_child * f for grand_child in child) for child in t), + ) + + +@pytest.mark.uses_tuple_args +def test_nested_tuple_comprehension_shadowing_names(cartesian_case): + @gtx.field_operator + def testee( + vals: tuple[tuple[cases.IField, ...], ...], factor: int32 + ) -> tuple[tuple[cases.IField, ...], ...]: + return tuple(tuple(child * factor for child in child) for child in vals) + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda t, f: tuple(tuple(child * f for child in child) for child in t), + ) + + +@pytest.mark.uses_tuple_args +def test_multi_target_tuple_comprehension(cartesian_case): + @gtx.field_operator + def testee(nested_tuple: tuple[tuple[int32, cases.IField], ...]) -> tuple[cases.IField, ...]: + return tuple(factor * tracer for factor, tracer in nested_tuple) + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda t: tuple(f * el for f, el in t), + ) + + @pytest.mark.uses_tuple_args @pytest.mark.uses_unstructured_shift @pytest.mark.xfail(reason="Iterator of tuple approach in lowering does not allow this.") @@ -254,5 +335,5 @@ def test_tuple_unpacking_too_few_values(cartesian_case): @gtx.field_operator(backend=cartesian_case.backend) def _invalid_unpack() -> tuple[int32, float64, int32]: - a, b, c = 1 + a, _b, _c = 1 return a diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_where.py index 7ea11a7b69..7b93d7892e 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_where.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_where.py @@ -7,16 +7,27 @@ # SPDX-License-Identifier: BSD-3-Clause import functools + import numpy as np -from typing import Tuple import pytest -from next_tests.integration_tests.cases import IDim, JDim, KDim, cartesian_case + from gt4py import next as gtx -from gt4py.next import float64, int32 -from gt4py.next.ffront.fbuiltins import where, broadcast +from gt4py.next import common, float64, int32, neighbor_sum +from gt4py.next.ffront.fbuiltins import broadcast, where + from next_tests.integration_tests import cases +from next_tests.integration_tests.cases import ( + E2V, + E2VDim, + IDim, + JDim, + KDim, + cartesian_case, + unstructured_case, +) from next_tests.integration_tests.cases_utils import ( exec_alloc_descriptor, + mesh_descriptor, ) @@ -113,6 +124,41 @@ def testee( ) +@pytest.mark.uses_tuple_returns +@pytest.mark.uses_unstructured_shift +def test_with_tuples_and_local_condition(unstructured_case): + @gtx.field_operator + def testee( + a: cases.VField, + b: cases.VField, + c: cases.VField, + d: cases.VField, + ) -> tuple[cases.EField, cases.EField]: + cond = a(E2V) > c(E2V) + t = where(cond, (a(E2V), b(E2V)), (c(E2V), d(E2V))) + return neighbor_sum(t[0], axis=E2VDim), neighbor_sum(t[1], axis=E2VDim) + + e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + cases.verify_with_default_data( + unstructured_case, + testee, + ref=lambda a, b, c, d: ( + np.sum( + np.where(a[e2v_table] > c[e2v_table], a[e2v_table], c[e2v_table]), + axis=1, + initial=0, + where=e2v_table != common._DEFAULT_SKIP_VALUE, + ), + np.sum( + np.where(a[e2v_table] > c[e2v_table], b[e2v_table], d[e2v_table]), + axis=1, + initial=0, + where=e2v_table != common._DEFAULT_SKIP_VALUE, + ), + ), + ) + + @pytest.mark.uses_tuple_returns def test_conditional_nested_tuple(cartesian_case): @gtx.field_operator @@ -124,7 +170,9 @@ def conditional_nested_tuple( return where(mask, ((a, b), (b, a)), ((5.0, 7.0), (7.0, 5.0))) size = cartesian_case.default_sizes[IDim] - mask = cartesian_case.as_field([IDim], np.random.choice(a=[False, True], size=size)) + mask = cartesian_case.as_field( + [IDim], np.random.default_rng().choice(a=[False, True], size=size) + ) a = cases.allocate(cartesian_case, conditional_nested_tuple, "a")() b = cases.allocate(cartesian_case, conditional_nested_tuple, "b")() @@ -158,7 +206,9 @@ def conditional( return where(mask, a, b) size = cartesian_case.default_sizes[IDim] - mask = cartesian_case.as_field([IDim], np.random.choice(a=[False, True], size=(size))) + mask = cartesian_case.as_field( + [IDim], np.random.default_rng().choice(a=[False, True], size=size) + ) a = cases.allocate(cartesian_case, conditional, "a")() b = cases.allocate(cartesian_case, conditional, "b")() out = cases.allocate(cartesian_case, conditional, cases.RETURN)() @@ -180,7 +230,9 @@ def conditional_promotion(mask: cases.IBoolField, a: cases.IFloatField) -> cases return where(mask, a, 10.0) size = cartesian_case.default_sizes[IDim] - mask = cartesian_case.as_field([IDim], np.random.choice(a=[False, True], size=(size))) + mask = cartesian_case.as_field( + [IDim], np.random.default_rng().choice(a=[False, True], size=size) + ) a = cases.allocate(cartesian_case, conditional_promotion, "a")() out = cases.allocate(cartesian_case, conditional_promotion, cases.RETURN)() ref = np.where(mask.asnumpy(), a.asnumpy(), 10.0) @@ -214,7 +266,9 @@ def conditional_program( conditional_shifted(mask, a, b, out=out) size = cartesian_case.default_sizes[IDim] + 1 - mask = cartesian_case.as_field([IDim], np.random.choice(a=[False, True], size=(size))) + mask = cartesian_case.as_field( + [IDim], np.random.default_rng().choice(a=[False, True], size=size) + ) a = cases.allocate(cartesian_case, conditional_program, "a").extend({IDim: (0, 1)})() b = cases.allocate(cartesian_case, conditional_program, "b").extend({IDim: (0, 1)})() out = cases.allocate(cartesian_case, conditional_shifted, cases.RETURN)() diff --git a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py index 62072cab18..50f03b0886 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py @@ -356,3 +356,38 @@ def broken(a: BrokenFieldAlias) -> gtx.Field[[IDim], float64]: assert err.location.line == err.location.end_line assert err.location.end_column - err.location.column == len("a: BrokenFieldAlias") assert re.search(r"\| +\^{19}(?!\^)", str(err)), str(err) + + +def test_bare_generator_expression_points_to_tuple_constructor(): + def with_genexp(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: + b = (x for x in (a, a)) # noqa: F841 [unused-variable] + return a + + err = parse_error(with_genexp) + + assert isinstance(err, errors.UnsupportedPythonFeatureError) + assert err.message == "Unsupported Python syntax: generator expression." + assert any("'tuple(...)'" in hint for hint in err.hints) + + +def test_nested_generator_expression_suggests_single_for_clause(): + def with_nested_genexp(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: + b = tuple(x for x in (a, a) for y in (a, a)) # noqa: F841 [unused-variable] + return a + + err = parse_error(with_nested_genexp) + + assert err.message == "Nested generator expressions are not supported." + assert any("single 'for' clause" in hint for hint in err.hints) + + +def test_conditional_in_generator_expression_explains_static_length(): + def with_filtered_genexp(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: + b = tuple(x for x in (a, a) if True) # noqa: F841 [unused-variable] + return a + + err = parse_error(with_filtered_genexp) + + assert err.message == "Conditionals are not supported in generator expressions." + assert any("known at compile time" in note for note in err.notes) + assert "if True" in str(err) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index 407088e4f8..9c4dd08266 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py @@ -42,7 +42,6 @@ from gt4py.next.type_system import type_specifications as ts, type_translation from gt4py.next.iterator import ir as itir - Edge = gtx.Dimension("Edge") Vertex = gtx.Dimension("Vertex") V2EDim = gtx.Dimension("V2E", gtx.DimensionKind.LOCAL) @@ -382,6 +381,184 @@ def foo( assert lowered_inlined.expr == reference +def test_fixed_len_tuple_comprehension(): + def foo(a: tuple[gtx.Field[[TDim], float64], gtx.Field[[TDim], float64]], factor: float64): + return tuple(el * factor for el in a) + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + lowered_inlined = inline_lambdas.InlineLambdas.apply(lowered) + + iterable = "__tuple_comprh_0" + reference = im.let(iterable, "a")( + im.make_tuple( + im.let("el", im.tuple_get(0, iterable))(im.op_as_fieldop("multiplies")("el", "factor")), + im.let("el", im.tuple_get(1, iterable))(im.op_as_fieldop("multiplies")("el", "factor")), + ) + ) + reference_inlined = inline_lambdas.InlineLambdas.apply(reference) + + assert lowered_inlined.expr == reference_inlined + + +def test_var_len_tuple_comprehension_scalar(): + def foo(a: tuple[float64, ...]): + return tuple(2.0 * el for el in a) + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + + reference = im.call( + im.call("map_tuple")(im.lambda_("el")(im.multiplies_(im.literal("2.0", "float64"), "el"))) + )("a") + + assert lowered.expr == reference + + +def test_var_len_tuple_comprehension_field(): + def foo(a: tuple[gtx.Field[[TDim], float64], ...], factor: float64): + return tuple(el * factor for el in a) + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + + reference = im.call( + im.call("map_tuple")(im.lambda_("el")(im.op_as_fieldop("multiplies")("el", "factor"))) + )("a") + + assert lowered.expr == reference + + +def test_var_len_tuple_comprehension_field_with_local_dim(): + def foo(a: tuple[gtx.Field[[Vertex, V2EDim], float64], ...]): + return tuple(2.0 * el for el in a) + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + + two = im.literal("2.0", "float64") + reference = im.call( + im.call("map_tuple")( + im.lambda_("el")( + im.op_as_fieldop(im.map_list("multiplies"))( + im.op_as_fieldop("make_const_list")(two), "el" + ) + ) + ) + )("a") + + assert lowered.expr == reference + + +def test_fixed_len_tuple_comprehension_local_field(): + def foo(a: gtx.Field[[Edge], float64]): + return tuple(2.0 * el for el in (a(V2E), a(V2E))) + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + lowered_inlined = inline_lambdas.InlineLambdas.apply(lowered) + + iterable = "__tuple_comprh_0" + two = im.literal("2.0", "float64") + reference = im.let( + iterable, + im.make_tuple(im.as_fieldop_neighbors("V2E", "a"), im.as_fieldop_neighbors("V2E", "a")), + )( + im.make_tuple( + im.let("el", im.tuple_get(0, iterable))( + im.op_as_fieldop(im.map_list("multiplies"))( + im.op_as_fieldop("make_const_list")(two), "el" + ) + ), + im.let("el", im.tuple_get(1, iterable))( + im.op_as_fieldop(im.map_list("multiplies"))( + im.op_as_fieldop("make_const_list")(two), "el" + ) + ), + ) + ) + reference_inlined = inline_lambdas.InlineLambdas.apply(reference) + + assert lowered_inlined.expr == reference_inlined + + +def test_fixed_len_tuple_comprehension_mixed_local_field_tuple_target(): + def foo( + a: gtx.Field[[Edge], float64], + b: gtx.Field[[Vertex], float64], + c: gtx.Field[[Edge], float64], + d: gtx.Field[[Vertex], float64], + ): + return tuple( + 2.0 * local_el + scalar_el for local_el, scalar_el in ((a(V2E), b), (c(V2E), d)) + ) + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + lowered_inlined = inline_lambdas.InlineLambdas.apply(lowered) + + iterable = "__tuple_comprh_0" + two = im.literal("2.0", "float64") + local_term = im.op_as_fieldop(im.map_list("multiplies"))( + im.op_as_fieldop("make_const_list")(two), "local_el" + ) + element_expr = im.op_as_fieldop(im.map_list("plus"))( + local_term, im.op_as_fieldop("make_const_list")("scalar_el") + ) + tuple_el_0 = im.tuple_get(0, iterable) + tuple_el_1 = im.tuple_get(1, iterable) + mapped_tuple1 = im.make_tuple(im.as_fieldop_neighbors("V2E", "a"), "b") + mapped_tuple2 = im.make_tuple(im.as_fieldop_neighbors("V2E", "c"), "d") + reference = im.let(iterable, im.make_tuple(mapped_tuple1, mapped_tuple2))( + im.make_tuple( + im.let( + ("local_el", im.tuple_get(0, tuple_el_0)), + ("scalar_el", im.tuple_get(1, tuple_el_0)), + )(element_expr), + im.let( + ("local_el", im.tuple_get(0, tuple_el_1)), + ("scalar_el", im.tuple_get(1, tuple_el_1)), + )(element_expr), + ) + ) + reference_inlined = inline_lambdas.InlineLambdas.apply(reference) + + assert lowered_inlined.expr == reference_inlined + + +def test_fixed_len_tuple_comprehension_mixed_local_field(): + def foo(a: gtx.Field[[Edge], float64], b: gtx.Field[[Vertex], float64]): + return tuple(2.0 * el for el in (a(V2E), b)) + + with pytest.raises( + NotImplementedError, + match="fixed-length tuples require all iterable elements to have the same type", + ): + FieldOperatorParser.apply_to_function(foo) + + +def test_fixed_len_tuple_comprehension_mixed_field_domains(): + def foo(a: gtx.Field[[Edge], float64], b: gtx.Field[[Vertex], float64]): + return tuple(2.0 * el for el in (a, b)) + + with pytest.raises( + NotImplementedError, + match="fixed-length tuples require all iterable elements to have the same type", + ): + FieldOperatorParser.apply_to_function(foo) + + +def test_fixed_len_tuple_comprehension_tuple_target_mixed_element_types(): + def foo(a: gtx.Field[[Edge], float64], b: gtx.Field[[Vertex], float64]): + return tuple(2.0 * left + right for left, right in ((a(V2E), b), (b, b))) + + with pytest.raises( + NotImplementedError, + match="fixed-length tuples require all iterable elements to have the same type", + ): + FieldOperatorParser.apply_to_function(foo) + + def test_unary_minus(): def foo(inp: gtx.Field[[TDim], float64]): return -inp diff --git a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py index a75e7bb031..bed07f0685 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py @@ -490,3 +490,36 @@ def tuple_index_failure( with pytest.raises(errors.DSLError, match=r"need .* literal"): _ = FieldOperatorParser.apply_to_function(tuple_index_failure) + + +def test_tuple_compr_non_tuple_iterable_failure(): + def testee(arg: float): + return tuple(_ for _ in arg) + + with pytest.raises( + errors.DSLError, + match=re.escape("Iterable in generator expression must be a tuple, got 'float64'."), + ): + _ = FieldOperatorParser.apply_to_function(testee) + + +def test_nested_tuple_compr_failure(): + def testee(nested_tuple: tuple[tuple[gtx.Field[[TDim], float64], ...], ...], factor: int32): + return tuple(grandchild * factor for child in nested_tuple for grandchild in child) + + with pytest.raises( + errors.DSLError, + match=re.escape("Nested generator expressions are not supported."), + ): + _ = FieldOperatorParser.apply_to_function(testee) + + +def test_tuple_compr_unpacking_failure(): + def testee(arg: tuple[int32, ...]): + return tuple(a * b for a, b in arg) + + with pytest.raises( + errors.DSLError, + match=re.escape("Cannot unpack non-iterable 'int32' object."), + ): + _ = FieldOperatorParser.apply_to_function(testee) diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py index 91d644dcac..fe259f18d0 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py @@ -216,15 +216,17 @@ def test_invalid_symbol_types(): type_translation.from_type_hint("foo") # Tuples - with pytest.raises(ValueError, match="least one argument"): - type_translation.from_type_hint(typing.Tuple) - with pytest.raises(ValueError, match="least one argument"): - type_translation.from_type_hint(tuple) - - with pytest.raises(ValueError, match="Unbound tuples"): - type_translation.from_type_hint(tuple[int, ...]) - with pytest.raises(ValueError, match="Unbound tuples"): - type_translation.from_type_hint(typing.Tuple["float", ...]) + # Bare `tuple` and `typing.Tuple` (unparameterized) both return a DeferredType. + assert type_translation.from_type_hint(typing.Tuple) == ts.DeferredType(constraint=ts.TupleType) + assert type_translation.from_type_hint(tuple) == ts.DeferredType(constraint=ts.TupleType) + + # Variadic tuples (`tuple[T, ...]`) are now valid — returns a VarArgType. + assert type_translation.from_type_hint(tuple[int, ...]) == ts.VarArgType( + element_type=ts.ScalarType(kind=ts.ScalarKind.INT64) + ) + assert type_translation.from_type_hint(typing.Tuple["float", ...]) == ts.VarArgType( + element_type=ts.ScalarType(kind=ts.ScalarKind.FLOAT64) + ) # Fields with pytest.raises(ValueError, match="Field type requires two arguments"): From 4a117afc563f4e6f0905a4178b521754a6ea2daa Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Thu, 27 Aug 2026 22:47:48 +0000 Subject: [PATCH 4/8] Address Fable review: builtin-constructor arity diagnostic and length-mismatch xfail test --- src/gt4py/next/ffront/func_to_foast.py | 5 ++++ .../feature_tests/ffront_tests/test_tuples.py | 25 +++++++++++++++++++ .../ffront_tests/test_func_to_foast.py | 14 +++++++++++ 3 files changed, 44 insertions(+) diff --git a/src/gt4py/next/ffront/func_to_foast.py b/src/gt4py/next/ffront/func_to_foast.py index d2926a90fc..b210027d5d 100644 --- a/src/gt4py/next/ffront/func_to_foast.py +++ b/src/gt4py/next/ffront/func_to_foast.py @@ -543,6 +543,11 @@ def visit_NotEq(self, node: ast.NotEq, **kwargs: Any) -> foast.CompareOperator: def _verify_builtin_type_constructor(self, node: ast.Call) -> None: assert isinstance(node.func, ast.Name) + if len(node.args) != 1: + raise errors.DSLError( + self.get_location(node), + f"'{self._func_name(node)}()' takes exactly one argument, got {len(node.args)}.", + ) (arg,) = node.args if not ( isinstance(arg, ast.Constant) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index 090f6a7e7e..392b900841 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -141,6 +141,31 @@ def testee(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IFie ) +@pytest.mark.uses_tuple_args +@pytest.mark.xfail( + strict=True, + reason="The consistency of the lengths of variable-length tuple arguments and the " + "'out' argument is not checked anywhere; the mismatch only surfaces as an internal " + "error (e.g. a bare 'AssertionError' from 'utils.tree_map') deep in the toolchain.", +) +def test_var_len_tuple_length_mismatch_rejected(cartesian_case): + # The return length equals the length of `a` or `b` depending on a runtime value, so + # with `len(a) != len(b)` no single length can be consistent with `out` — this should + # be rejected with a proper diagnostic when the concrete argument types are known. + @gtx.field_operator + def testee( + a: tuple[cases.IField, ...], b: tuple[cases.IField, ...], c: bool + ) -> tuple[cases.IField, ...]: + return a if c else b + + a = cases.allocate(cartesian_case, testee, "a")()[:2] + b = cases.allocate(cartesian_case, testee, "b")() # allocated length: 3 + out = cases.allocate(cartesian_case, testee, cases.RETURN)()[:2] + + with pytest.raises(errors.DSLError): + cases.run(cartesian_case, testee, a, b, False, out=out) + + @pytest.mark.uses_tuple_args def test_tuple_comprehension_other_fo(cartesian_case): @gtx.field_operator diff --git a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py index bed07f0685..b72e567a12 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py @@ -269,6 +269,20 @@ def cast_scalar_temp(): _ = FieldOperatorParser.apply_to_function(cast_scalar_temp) +def test_scalar_cast_wrong_arity(): + def cast_scalar_two_args(): + return int32(1, 2) + + with pytest.raises(errors.DSLError, match=r"takes exactly one argument, got 2"): + _ = FieldOperatorParser.apply_to_function(cast_scalar_two_args) + + def cast_scalar_no_args(): + return int32() + + with pytest.raises(errors.DSLError, match=r"takes exactly one argument, got 0"): + _ = FieldOperatorParser.apply_to_function(cast_scalar_no_args) + + def test_conditional_wrong_mask_type(): def conditional_wrong_mask_type(a: gtx.Field[[TDim], float64]) -> gtx.Field[[TDim], float64]: return where(a, a, a) From 666c8b4df8e687d0bb0a14025a66ce5074b2e0b8 Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Tue, 1 Sep 2026 07:34:02 +0000 Subject: [PATCH 5/8] Address review comments (unpack diagnostics, DSLError for heterogeneity, tuple[()] handling, empty-vararg test) --- .../ffront/foast_passes/type_deduction.py | 32 +++++++++++-------- src/gt4py/next/type_system/type_info.py | 2 ++ .../next/type_system/type_translation.py | 15 ++++++--- .../ffront_tests/test_foast_to_gtir.py | 12 +++---- .../ffront_tests/test_type_deduction.py | 18 +++++++++++ .../type_system_tests/test_type_info.py | 12 +++++++ .../test_type_translation.py | 6 ++++ 7 files changed, 73 insertions(+), 24 deletions(-) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 9e52486921..59215b8ef2 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -760,17 +760,20 @@ def _deduce_tuple_comprehension_target_type( @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_, **kwargs) - except IndexError: - raise errors.DSLError( - target_el.location, f"Cannot unpack non-iterable '{type_}' object." - ) from None + type_ = element_type + for i in path: + if not isinstance(type_, ts.TupleType): + raise errors.DSLError( + target_el.location, f"Cannot unpack non-iterable '{type_}' object." + ) + if len(type_.types) <= i: + raise errors.DSLError( + target_el.location, + f"Not enough values to unpack (expected at least {i + 1}, " + f"got {len(type_.types)}).", + ) + type_ = type_.types[i] + return self.visit(target_el, refine_type=type_, **kwargs) return process_target(target) @@ -822,9 +825,10 @@ def visit_TupleComprehension( element_types = cast(list[ts.DataType], iterable.type.types) # Only homogeneous iterables are supported, see ADR 0028. 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 (see ADR 0028)." + raise errors.DSLError( + iterable.location, + "Tuple comprehensions over fixed-length tuples with differently typed " + "iterable elements are not implemented (see ADR 0028).", ) new_mapper = self._deduce_tuple_comprehension_mapper( node, target, element_types[0], **kwargs diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index e3651b1c62..f4f5288770 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -553,6 +553,8 @@ def is_concretizable(symbol_type: ts.TypeSpec, to_type: ts.TypeSpec) -> bool: if isinstance(symbol_type, ts.VarArgType) and isinstance(to_type, ts.VarArgType): return is_concretizable(symbol_type.element_type, to_type.element_type) if isinstance(symbol_type, ts.VarArgType) and isinstance(to_type, ts.TupleType): + # The empty tuple is a valid concretization of any variable-length tuple: there + # are no elements whose type could conflict with the element type. if len(to_type.types) == 0 or ( all(type_ == to_type.types[0] for type_ in to_type.types) and is_concretizable(symbol_type.element_type, to_type.types[0]) diff --git a/src/gt4py/next/type_system/type_translation.py b/src/gt4py/next/type_system/type_translation.py index 4ff4d3970b..13afc6a313 100644 --- a/src/gt4py/next/type_system/type_translation.py +++ b/src/gt4py/next/type_system/type_translation.py @@ -188,8 +188,14 @@ def canonicalize_type_hint( type_hint = _resolve_type_alias(type_hint) canonical_type = typing.get_origin(type_hint) or type_hint - # In order to distinguish `tuple` from `tuple[()]`, the former returns None here. - args = typing.get_args(type_hint) if typing.get_origin(type_hint) else None + # `None` (rather than `()`) marks an unparametrized annotation, distinguishing bare + # `tuple` / `typing.Tuple` from the empty-tuple annotation `tuple[()]`. Unsubscripted + # `typing` aliases have a truthy origin but no `__args__` attribute. + args = ( + typing.get_args(type_hint) + if typing.get_origin(type_hint) and hasattr(type_hint, "__args__") + else None + ) return canonical_type, args, tuple(extra_args) @@ -222,8 +228,9 @@ def from_type_hint( # Variable-length tuple, e.g. `tuple[int32, ...]` elif isinstance(args, tuple) and len(args) == 2 and args[1] is Ellipsis: return ts.VarArgType(element_type=from_type_hint_same_ns(args[0])) - # Unparametrized annotation, i.e. `tuple` or `tuple[()]` - elif args is None or (isinstance(args, tuple) and len(args) == 0): + # Unparametrized annotation, i.e. bare `tuple` / `typing.Tuple` (the + # empty-tuple annotation `tuple[()]` has `args == ()` and is rejected below) + elif args is None: # TODO(tehrengruber): We use `DeferredType` until we have an actual representation # for a generic type. return ts.DeferredType(constraint=ts.TupleType) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index 9c4dd08266..e6d7249c36 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py @@ -531,8 +531,8 @@ def foo(a: gtx.Field[[Edge], float64], b: gtx.Field[[Vertex], float64]): return tuple(2.0 * el for el in (a(V2E), b)) with pytest.raises( - NotImplementedError, - match="fixed-length tuples require all iterable elements to have the same type", + errors.DSLError, + match="fixed-length tuples with differently typed iterable elements are not implemented", ): FieldOperatorParser.apply_to_function(foo) @@ -542,8 +542,8 @@ def foo(a: gtx.Field[[Edge], float64], b: gtx.Field[[Vertex], float64]): return tuple(2.0 * el for el in (a, b)) with pytest.raises( - NotImplementedError, - match="fixed-length tuples require all iterable elements to have the same type", + errors.DSLError, + match="fixed-length tuples with differently typed iterable elements are not implemented", ): FieldOperatorParser.apply_to_function(foo) @@ -553,8 +553,8 @@ def foo(a: gtx.Field[[Edge], float64], b: gtx.Field[[Vertex], float64]): return tuple(2.0 * left + right for left, right in ((a(V2E), b), (b, b))) with pytest.raises( - NotImplementedError, - match="fixed-length tuples require all iterable elements to have the same type", + errors.DSLError, + match="fixed-length tuples with differently typed iterable elements are not implemented", ): FieldOperatorParser.apply_to_function(foo) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py index 22bd1a7a9e..ef1ddb3a3d 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py @@ -74,6 +74,24 @@ def temp_tuple(a: Field[[TDim], float64], b: Field[[TDim], int64]): ) +def test_tuple_comprehension_unpack_non_tuple(): + def foo(it: tuple[Field[[TDim], float64], ...]): + return tuple(a + b for a, b in it) + + with pytest.raises(errors.DSLError, match=r"Cannot unpack non-iterable"): + _ = FieldOperatorParser.apply_to_function(foo) + + +def test_tuple_comprehension_unpack_too_short_tuple(): + def foo(it: tuple[tuple[Field[[TDim], float64], Field[[TDim], float64]], ...]): + return tuple(a + b + c for a, b, c in it) + + with pytest.raises( + errors.DSLError, match=r"Not enough values to unpack \(expected at least 3, got 2\)" + ): + _ = FieldOperatorParser.apply_to_function(foo) + + def test_adding_bool(): """Expect an error when using arithmetic on bools.""" diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py index 35c3d2eba1..271688d655 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py @@ -462,3 +462,15 @@ def test_return_type( ) def test_needs_value_extraction(type_spec: ts.TypeSpec, expected: bool): assert type_info.needs_value_extraction(type_spec) is expected + + +def test_is_concretizable_vararg_to_tuple(): + float_type = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) + int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) + vararg = ts.VarArgType(element_type=float_type) + + assert type_info.is_concretizable(vararg, ts.TupleType(types=[float_type, float_type])) + assert not type_info.is_concretizable(vararg, ts.TupleType(types=[int_type])) + # The empty tuple is a valid concretization of any variable-length tuple: there are + # no elements whose type could conflict with the element type. + assert type_info.is_concretizable(vararg, ts.TupleType(types=[])) diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py index fe259f18d0..6e55ca0614 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py @@ -220,6 +220,12 @@ def test_invalid_symbol_types(): assert type_translation.from_type_hint(typing.Tuple) == ts.DeferredType(constraint=ts.TupleType) assert type_translation.from_type_hint(tuple) == ts.DeferredType(constraint=ts.TupleType) + # The empty-tuple annotation is not a valid type annotation. + with pytest.raises(ValueError, match="Tuple annotation"): + type_translation.from_type_hint(tuple[()]) + with pytest.raises(ValueError, match="Tuple annotation"): + type_translation.from_type_hint(typing.Tuple[()]) + # Variadic tuples (`tuple[T, ...]`) are now valid — returns a VarArgType. assert type_translation.from_type_hint(tuple[int, ...]) == ts.VarArgType( element_type=ts.ScalarType(kind=ts.ScalarKind.INT64) From 0f11a59b55453e21f0a30e794d5b290b3482e12f Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Tue, 1 Sep 2026 09:10:54 +0000 Subject: [PATCH 6/8] Add empty-tuple comprehension test, xfailed on compiled backends --- .../feature_tests/ffront_tests/test_tuples.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index 392b900841..39c53e0073 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -141,6 +141,22 @@ def testee(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IFie ) +@pytest.mark.uses_tuple_args +def test_var_len_tuple_comprehension_empty(cartesian_case): + # The empty tuple is a valid concretization of a variable-length tuple; the program + # then has nothing to write and must simply run without error. + @gtx.field_operator + def testee(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IField, ...]: + return tuple(tracer * factor for tracer in tracers) + + if cartesian_case.backend is not None: + pytest.xfail( + "The compiled toolchain cannot execute a program with an empty tuple output: " + "a 'SetAt' for an empty 'out' has no target and no domain to derive." + ) + cases.run(cartesian_case, testee, (), 42, out=()) + + @pytest.mark.uses_tuple_args @pytest.mark.xfail( strict=True, From 681ce77a1a27081c339dde3a66ae6c8e0b2de738 Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Tue, 1 Sep 2026 09:20:55 +0000 Subject: [PATCH 7/8] Add unit test for the generic-tuple constructor type --- .../type_system_tests/test_type_translation.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py index 6e55ca0614..ace88bee3e 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py @@ -210,6 +210,23 @@ def test_make_symbol_type_from_typing(value, expected): assert type_translation.from_type_hint(value) == expected +def test_generic_tuple_constructor_type(): + # The type of the `tuple` builtin itself: `type[tuple]` (e.g. the symbol `tuple` + # referenced inside a field operator) maps to a constructor for a generic tuple. + expected = ts.ConstructorType( + definition=ts.FunctionType( + pos_only_args=[ts.DeferredType(constraint=None)], + pos_or_kw_args={}, + kw_only_args={}, + returns=ts.DeferredType(constraint=ts.VarArgType), + ) + ) + assert ( + type_translation.make_constructor_type(ts.DeferredType(constraint=ts.TupleType)) == expected + ) + assert type_translation.from_type_hint(type[tuple]) == expected + + def test_invalid_symbol_types(): # Forward references with pytest.raises(ValueError, match="undefined forward references"): From 5da530ef72fe9306a7e2293b8acd7fd3df67dfef Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 11 Sep 2026 12:06:49 +0200 Subject: [PATCH 8/8] test[next]: improve coverage for closure vars in tuple comprehension (#2868) --- .../feature_tests/ffront_tests/test_tuples.py | 48 ++++++++++++++++++- .../ffront_tests/test_source_utils.py | 31 ++++++++++-- 2 files changed, 73 insertions(+), 6 deletions(-) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index 39c53e0073..5a1b38cab5 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -10,7 +10,7 @@ import pytest import gt4py.next as gtx -from gt4py.next import broadcast, errors, float64, int32, neighbor_sum, utils as gt_utils +from gt4py.next import broadcast, common, errors, float64, int32, neighbor_sum, utils as gt_utils from next_tests.integration_tests import cases from next_tests.integration_tests.cases import ( @@ -199,6 +199,52 @@ def testee(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IFie ) +@gtx.field_operator +def _module_level_inner(tracer: cases.IField, factor: int32) -> cases.IField: + return tracer * factor + + +@pytest.mark.uses_tuple_args +def test_tuple_comprehension_module_level_fo(cartesian_case): + @gtx.field_operator + def testee(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IField, ...]: + return tuple(_module_level_inner(tracer, factor) for tracer in tracers) + + cases.verify_with_default_data( + cartesian_case, testee, ref=lambda t, f: tuple(el * f for el in t) + ) + + +@pytest.mark.uses_tuple_args +@pytest.mark.uses_unstructured_shift +def test_tuple_comprehension_module_level_builtin(unstructured_case): + @gtx.field_operator + def testee(tracers: tuple[cases.EField, ...]) -> tuple[cases.VField, ...]: + return tuple(neighbor_sum(tracer(V2E), axis=V2EDim) for tracer in tracers) + + v2e = unstructured_case.offset_provider["V2E"].asnumpy() + valid = v2e != common._DEFAULT_SKIP_VALUE + cases.verify_with_default_data( + unstructured_case, + testee, + ref=lambda t: tuple(np.sum(el[v2e], axis=1, where=valid) for el in t), + ) + + +@pytest.mark.uses_tuple_args +def test_tuple_comprehension_module_level_fo_used_outside_too(cartesian_case): + @gtx.field_operator + def testee( + tracers: tuple[cases.IField, ...], a: cases.IField, factor: int32 + ) -> tuple[cases.IField, ...]: + unrelated = _module_level_inner(a, factor) + return tuple(_module_level_inner(tracer, factor) for tracer in tracers) + + cases.verify_with_default_data( + cartesian_case, testee, ref=lambda t, a, f: tuple(el * f for el in t) + ) + + @pytest.mark.uses_tuple_args def test_nested_tuple_comprehension(cartesian_case): @gtx.field_operator diff --git a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py index 5735b4c25b..9a08544d91 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py @@ -10,17 +10,19 @@ A generator expression, lambda or nested function compiles to its own code object, so names referenced only inside it appear in that object's `co_names` and never in -the enclosing function's. Collecting closure variables from the enclosing function -alone therefore misses them; `get_closure_vars_from_function` takes the global names -from the compiler's symbol table of the source instead, which covers nested scopes. +the enclosing function's. `get_closure_vars_from_function` therefore takes the global +names from the compiler's symbol table of the source, which covers nested scopes. Free variables need no such treatment: the enclosing code object carries a cell for them. """ +import pytest + import gt4py.next as gtx from gt4py.next import Dims, Dimension, float64, neighbor_sum from gt4py.next.ffront import source_utils +from gt4py.next.ffront.func_to_foast import FieldOperatorParser from gt4py.next.ffront.source_utils import get_closure_vars_from_function @@ -46,6 +48,27 @@ def _module_level_operator(tracers: tuple[CField, ...], factor: float64) -> tupl return tuple(scale(t, factor) for t in tracers) +@pytest.mark.parametrize( + "func", [_builtin_and_offset, _module_level_operator], ids=["builtin_and_offset", "operator"] +) +def test_module_level_name_used_only_in_comprehension(func): + """Module-level names are reachable from inside a comprehension body.""" + parsed = FieldOperatorParser.apply_to_function(func) + assert parsed.type is not None + + +def test_free_variables_still_resolve_through_the_parser(): + @gtx.field_operator + def local_scale(f: CField, factor: float64) -> CField: + return f * factor + + def uses_freevar(tracers: tuple[CField, ...], factor: float64) -> tuple[CField, ...]: + return tuple(local_scale(t, factor) for t in tracers) + + parsed = FieldOperatorParser.apply_to_function(uses_freevar) + assert parsed.type is not None + + def test_names_are_collected_from_the_nested_code_object(): """The names live in the generator expression's code object, not the function's.""" assert "neighbor_sum" not in _builtin_and_offset.__code__.co_names @@ -62,8 +85,6 @@ def test_comprehension_target_is_not_collected(): def test_free_variables_still_resolve(): - """The path that already worked, kept as a guard.""" - @gtx.field_operator def local_scale(f: CField, factor: float64) -> CField: return f * factor