diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index 3b9f97592f..b88c7d8f93 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -31,6 +31,7 @@ Field, GridType, UnitRange, + XTuple, as_non_staggered, domain, flip_staggered, @@ -118,6 +119,7 @@ "DimensionKind", "Dims", "Field", + "XTuple", "CartesianConnectivity", "Connectivity", "GridType", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 741cbeda33..6d364fba26 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -57,6 +57,87 @@ class Dims(tuple[Unpack[ShapeTs]]): ... +class XTuple(tuple[Unpack[ShapeTs]]): + """ + Tuple on which binary arithmetic and logical operators apply element-wise. + + Non-tuple operands (e.g. scalars, fields) are broadcast against the tuple structure. + Nested element-wise behavior requires nested `XTuple`s; plain `tuple` operands are + rejected, mirroring the type deduction rules of the DSL. + """ + + def _elementwise_op(self, other: Any, op: Callable[[Any, Any], Any]) -> XTuple: + self_elems: tuple[Any, ...] = tuple(self) + if isinstance(other, XTuple): + other_elems: tuple[Any, ...] = tuple(other) + if len(self_elems) != len(other_elems): + raise ValueError( + f"Element-wise operations require 'XTuple's of equal length, " + f"got {len(self_elems)} and {len(other_elems)}." + ) + return XTuple(op(el, other_el) for el, other_el in zip(self_elems, other_elems)) + if isinstance(other, tuple): + raise TypeError("Element-wise operations require 'XTuple' operands, got 'tuple'.") + return XTuple(op(el, other) for el in self_elems) + + def __add__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a + b) + + def __radd__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b + a) + + def __sub__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a - b) + + def __rsub__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b - a) + + def __mul__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a * b) + + def __rmul__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b * a) + + def __truediv__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a / b) + + def __rtruediv__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b / a) + + def __floordiv__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a // b) + + def __rfloordiv__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b // a) + + def __mod__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a % b) + + def __rmod__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b % a) + + def __pow__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a**b) + + def __and__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a & b) + + def __rand__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b & a) + + def __or__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a | b) + + def __ror__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b | a) + + def __xor__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: a ^ b) + + def __rxor__(self, other: Any) -> XTuple: + return self._elementwise_op(other, lambda a, b: b ^ a) + + DimsT = TypeVar("DimsT", bound=Dims, covariant=True) Tag: TypeAlias = str diff --git a/src/gt4py/next/ffront/field_operator_ast.py b/src/gt4py/next/ffront/field_operator_ast.py index f84dd53d62..0b00f43dc2 100644 --- a/src/gt4py/next/ffront/field_operator_ast.py +++ b/src/gt4py/next/ffront/field_operator_ast.py @@ -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 diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index c9f51ad080..2a4bdca141 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 NestedTuple 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) @@ -456,6 +458,17 @@ 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): + try: + # The length of a variable-length tuple is only known when the concrete + # arguments are available, so the index can not be bounds-checked here. + foast_utils.expr_to_index(node.index) + except ValueError as ex: + raise errors.DSLError( + node.location, + f"Tuples need to be indexed with literal integers, got '{node.index}'.", + ) from ex + new_type = element_type case ts.OffsetType(source=source, target=(target1, target2)): if not target2.kind == DimensionKind.LOCAL: raise errors.DSLError( @@ -623,6 +636,11 @@ def _deduce_binop_type( err_msg = f"Unsupported operand type(s) for {node.op}: '{left.type}' and '{right.type}'." + if isinstance(left.type, (ts.XTupleType, ts.XVarArgType)) or isinstance( + right.type, (ts.XTupleType, ts.XVarArgType) + ): + return self._deduce_elementwise_binop_type(node, left=left, right=right) + if isinstance(left.type, (ts.ScalarType, ts.FieldType)) and isinstance( right.type, (ts.ScalarType, ts.FieldType) ): @@ -710,6 +728,110 @@ def _deduce_binop_type( else: raise errors.DSLError(node.location, err_msg) + def _deduce_elementwise_binop_type( + self, node: foast.BinOp, *, left: foast.Expr, right: foast.Expr + ) -> ts.XTupleType | ts.XVarArgType: + def operand_with_type(operand: foast.Expr, type_: ts.TypeSpec) -> foast.Expr: + return foast.Constant(value=None, location=operand.location, type=type_) + + def tuple_element_type( + type_: ts.TypeSpec, + ) -> ts.DataType | ts.DimensionType | ts.DeferredType: + if not isinstance(type_, (ts.DataType, ts.DimensionType, ts.DeferredType)): + raise errors.DSLError( + node.location, + f"Element-wise operator '{node.op}' produced unsupported tuple element type " + f"'{type_}'.", + ) + return type_ + + def vararg_element_type(type_: ts.TypeSpec) -> ts.DataType: + if not isinstance(type_, ts.DataType): + raise errors.DSLError( + node.location, + f"Element-wise operator '{node.op}' produced unsupported variadic tuple " + f"element type '{type_}'.", + ) + return type_ + + def deduce(left_type: ts.TypeSpec, right_type: ts.TypeSpec) -> ts.TypeSpec: + for type_ in (left_type, right_type): + if isinstance(type_, (ts.TupleType, ts.VarArgType)) and not isinstance( + type_, (ts.XTupleType, ts.XVarArgType) + ): + raise errors.DSLError( + node.location, + f"Element-wise operator '{node.op}' requires tuple operands to be 'XTuple', " + f"got '{type_}'.", + ) + + if isinstance(left_type, ts.XTupleType) and isinstance(right_type, ts.XTupleType): + if len(left_type.types) != len(right_type.types): + raise errors.DSLError( + node.location, + f"Element-wise operator '{node.op}' requires tuple operands to have the " + f"same structure, got '{left.type}' and '{right.type}'.", + ) + return ts.XTupleType( + types=[ + tuple_element_type(deduce(left_el_type, right_el_type)) + for left_el_type, right_el_type in zip( + left_type.types, right_type.types, strict=True + ) + ] + ) + + elif ( + isinstance(left_type, ts.XTupleType) and isinstance(right_type, ts.XVarArgType) + ) or (isinstance(left_type, ts.XVarArgType) and isinstance(right_type, ts.XTupleType)): + raise errors.DSLError( + node.location, + f"Element-wise operator '{node.op}' can not be applied between a fixed-length " + f"tuple and a variable-length tuple: '{left.type}' and '{right.type}'.", + ) + + elif isinstance(left_type, ts.XTupleType): + return ts.XTupleType( + types=[ + tuple_element_type(deduce(left_el_type, right_type)) + for left_el_type in left_type.types + ] + ) + elif isinstance(right_type, ts.XTupleType): + return ts.XTupleType( + types=[ + tuple_element_type(deduce(left_type, right_el_type)) + for right_el_type in right_type.types + ] + ) + + elif isinstance(left_type, ts.XVarArgType) and isinstance(right_type, ts.XVarArgType): + raise errors.DSLError( + node.location, + f"Element-wise operator '{node.op}' can not be applied between two " + f"variable-length tuples: '{left.type}' and '{right.type}'.", + ) + elif isinstance(left_type, ts.XVarArgType): + return ts.XVarArgType( + element_type=vararg_element_type(deduce(left_type.element_type, right_type)) + ) + elif isinstance(right_type, ts.XVarArgType): + return ts.XVarArgType( + element_type=vararg_element_type(deduce(left_type, right_type.element_type)) + ) + + result_type = self._deduce_binop_type( + node, + left=operand_with_type(left, left_type), + right=operand_with_type(right, right_type), + ) + assert result_type is not None + return result_type + + result = deduce(left.type, right.type) + assert isinstance(result, (ts.XTupleType, ts.XVarArgType)) + return result + def _check_operand_dtypes_match( self, node: foast.BinOp | foast.Compare, left: foast.Expr, right: foast.Expr ) -> None: @@ -743,6 +865,96 @@ 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]) + tuple_type_cls = ( + ts.XTupleType if isinstance(iterable.type, ts.XTupleType) else ts.TupleType + ) + result = foast.TupleComprehension( + inner=new_mapper, + iterable=iterable, + location=node.location, + type=tuple_type_cls(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 = deduce_mapper(element_type) + element_expr = new_mapper.element_expr + vararg_type_cls = ( + ts.XVarArgType if isinstance(iterable.type, ts.XVarArgType) else ts.VarArgType + ) + return_type = vararg_type_cls(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) @@ -1026,7 +1238,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 diff --git a/src/gt4py/next/ffront/foast_pretty_printer.py b/src/gt4py/next/ffront/foast_pretty_printer.py index 8b2e369501..8c4203be6d 100644 --- a/src/gt4py/next/ffront/foast_pretty_printer.py +++ b/src/gt4py/next/ffront/foast_pretty_printer.py @@ -120,6 +120,22 @@ 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: + def format_target(target: Any) -> str: + if isinstance(target, tuple): + return f"({', '.join(format_target(el) for el in target)})" + return self.visit(target, **kwargs) + + element_expr = self.visit(node.element_expr, **kwargs) + return f"{element_expr} for {format_target(node.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..dfbd6bf953 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -8,6 +8,8 @@ import dataclasses +import functools +import warnings from typing import Any, Callable, Optional from gt4py import eve @@ -259,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) @@ -301,8 +370,23 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: new_index = constant_folding.ConstantFolding.apply(self.visit(index, **kwargs)) assert isinstance(new_index, itir.Literal) assert isinstance(offset_name.type, ts.OffsetType) + if fbuiltins.is_cartesian_offset(offset_name.type): + # Deprecated: Cartesian shift via the subscript syntax `field(Off[i])`. + # We deduce the dimension from the offset type and emit a self-describing + # `CartesianOffset` (cf. the `Dim + idx` and `as_offset` cases). + warnings.warn( + f"Cartesian shifts via the subscript syntax 'field({offset_name.id}[i])' " + f"are deprecated; use 'field({offset_name.type.source.value} + i)' instead.", + DeprecationWarning, + stacklevel=2, + ) + dim = offset_name.type.source + shift_offset: itir.CartesianOffset | str = im.cartesian_offset(dim) + else: + # Unstructured neighbor selection, resolved through the offset provider. + shift_offset = offset_name.id current_expr = im.as_fieldop( - im.lambda_("__it")(im.deref(im.shift(offset_name.id, new_index)("__it"))) + im.lambda_("__it")(im.deref(im.shift(shift_offset, new_index)("__it"))) )(current_expr) # `field(Dim + idx)` (where `idx` is integer or half integer) case foast.BinOp( @@ -530,18 +614,96 @@ def visit_Constant(self, node: foast.Constant, **kwargs: Any) -> itir.Expr: def _lower_and_map(self, op: itir.Lambda | str, *args: Any, **kwargs: Any) -> itir.FunCall: return _map( - op, tuple(self.visit(arg, **kwargs) for arg in args), tuple(arg.type for arg in args) + op, + tuple(self.visit(arg, **kwargs) for arg in args), + tuple(arg.type for arg in args), + uids=self.uid_generator, ) +def _map_elementwise( + op: itir.Lambda | str, + lowered_args: tuple, + original_arg_types: tuple[ts.TypeSpec, ...], + uids: utils.IDGeneratorPool, +) -> itir.FunCall: + """ + Apply `op` element-wise over 'XTuple' arguments, broadcasting non-tuple arguments. + + Fixed-length tuple arguments are expanded directly, such that the tuple structure ends up + outside of the stencils, i.e. a tuple of `as_fieldop`s. Variable-length tuple arguments are + lowered to the `map_tuple` builtin, which is expanded into `make_tuple` by 'ExpandTupleMaps' + once the concrete length is known. + """ + if any(isinstance(t, ts.XTupleType) for t in original_arg_types): + # All tuple arguments are fixed-length tuples of equal length (ensured by type deduction). + lengths = {len(t.types) for t in original_arg_types if isinstance(t, ts.XTupleType)} + assert len(lengths) == 1 and not any( + isinstance(t, ts.XVarArgType) for t in original_arg_types + ) + (length,) = lengths + + # Let-bind the tuple arguments to avoid duplicating their expressions per element. + bindings: list[tuple[str, itir.Expr]] = [] + bound_args: list[itir.Expr] = [] + for arg, arg_type in zip(lowered_args, original_arg_types): + if isinstance(arg_type, ts.XTupleType): + name = next(uids["__elementwise"]) + bindings.append((name, arg)) + bound_args.append(im.ref(name)) + else: + bound_args.append(arg) + + elements = [ + _map( + op, + tuple( + im.tuple_get(i, arg) if isinstance(arg_type, ts.XTupleType) else arg + for arg, arg_type in zip(bound_args, original_arg_types) + ), + tuple( + arg_type.types[i] if isinstance(arg_type, ts.XTupleType) else arg_type + for arg_type in original_arg_types + ), + uids=uids, + ) + for i in range(length) + ] + return im.let(*bindings)(im.make_tuple(*elements)) + + # Exactly one variable-length tuple argument (ensured by type deduction). + (vararg_index,) = (i for i, t in enumerate(original_arg_types) if isinstance(t, ts.XVarArgType)) + vararg_type = original_arg_types[vararg_index] + assert isinstance(vararg_type, ts.XVarArgType) + param = next(uids["__elementwise"]) + body = _map( + op, + tuple(im.ref(param) if i == vararg_index else arg for i, arg in enumerate(lowered_args)), + tuple( + vararg_type.element_type if i == vararg_index else t + for i, t in enumerate(original_arg_types) + ), + uids=uids, + ) + return im.call(im.call("map_tuple")(im.lambda_(param)(body)))(lowered_args[vararg_index]) + + def _map( op: itir.Lambda | str, lowered_args: tuple, original_arg_types: tuple[ts.TypeSpec, ...], + uids: Optional[utils.IDGeneratorPool] = None, ) -> itir.FunCall: """ Mapping includes making the operation an `as_fieldop` (first kind of mapping), but also `itir.map_list`ing lists. + + Element-wise operations on 'XTuple' arguments are expanded such that the tuple structure ends + up outside of the stencils (see `_map_elementwise`). """ + if any(isinstance(t, (ts.XTupleType, ts.XVarArgType)) for t in original_arg_types): + assert uids is not None + return _map_elementwise(op, lowered_args, original_arg_types, uids) + if all( isinstance(t, (ts.ScalarType, ts.DimensionType, ts.DomainType)) for arg_type in original_arg_types diff --git a/src/gt4py/next/ffront/foast_to_past.py b/src/gt4py/next/ffront/foast_to_past.py index 9a560b7ff8..dc74f2828e 100644 --- a/src/gt4py/next/ffront/foast_to_past.py +++ b/src/gt4py/next/ffront/foast_to_past.py @@ -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] = [ diff --git a/src/gt4py/next/ffront/func_to_foast.py b/src/gt4py/next/ffront/func_to_foast.py index 14dceb25d1..6e7eb77236 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, NestedTuple 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 ( @@ -401,8 +401,13 @@ def visit_Return(self, node: ast.Return, **kwargs: Any) -> foast.Return: def visit_Expr(self, node: ast.Expr) -> foast.Expr: return self.visit(node.value) - def visit_Name(self, node: ast.Name, **kwargs: Any) -> foast.Name: - return foast.Name(id=node.id, location=self.get_location(node)) + def visit_Name(self, node: ast.Name, **kwargs: Any) -> foast.DataSymbol | foast.Name: + loc = self.get_location(node) + if isinstance(node.ctx, ast.Store): + return foast.DataSymbol(id=node.id, location=loc, type=ts.DeferredType(constraint=None)) + else: + assert isinstance(node.ctx, ast.Load) + return foast.Name(id=node.id, location=loc) def visit_UnaryOp(self, node: ast.UnaryOp, **kwargs: Any) -> foast.UnaryOp: return foast.UnaryOp( @@ -542,24 +547,66 @@ 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) + 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) + or (isinstance(arg, ast.UnaryOp) and isinstance(arg.operand, ast.Constant)) + or (node.func.id == "tuple" and isinstance(arg, ast.GeneratorExp)) + ): + raise errors.DSLError( + self.get_location(node), + f"'{self._func_name(node)}()' only takes literal arguments or a generator expression.", + ) 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(node), + "Nested generator expressions are not supported.", + ) + if gen_expr.generators[0].ifs != []: + raise errors.DSLError( + self.get_location(node), + "Conditionals are not supported in generator expressions as the size of " + "the result can only be deduced at runtime.", + ) + + def parse_target(target: ast.expr) -> NestedTuple[foast.DataSymbol]: + if isinstance(target, ast.Tuple): + return tuple(parse_target(el) for el in target.elts) + assert isinstance(target, ast.Name) + return self.visit(target, **kwargs) + + 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..530d407459 100644 --- a/src/gt4py/next/ffront/past_passes/type_deduction.py +++ b/src/gt4py/next/ffront/past_passes/type_deduction.py @@ -248,7 +248,7 @@ 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: + 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/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index a3c98a96fb..cb97cd6cc8 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -287,7 +287,7 @@ def deduce_return_type(tb: ts.FieldType | ts.ScalarType, fb: ts.FieldType | ts.S dtype = tb_dtype return_dims = common.promote_dims( - domain.dims, type_info.extract_dims(type_info.promote(tb, fb)) + domain.dims, type_info.extract_dims(tb), type_info.extract_dims(fb) ) return_type = ts.FieldType(dims=return_dims, dtype=dtype) return return_type @@ -401,10 +401,12 @@ def _canonicalize_nb_fields( def _canonicalize_nb_fields( - input_: ts.ScalarType - | ts.FieldType - | ts.TupleType - | tuple[ts.ScalarType | ts.FieldType | ts.TupleType, ...], + input_: ( + ts.ScalarType + | ts.FieldType + | ts.TupleType + | tuple[ts.ScalarType | ts.FieldType | ts.TupleType, ...] + ), ) -> ts.ScalarType | ts.FieldType | ts.TupleType: """ Transform neighbor / sparse field type by removal of local dimension and addition of corresponding `ListType` dtype. diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index 50f372f5a9..3fb0b26e50 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -507,13 +507,23 @@ def _is_generic(self) -> bool: Is the operator or program generic in the sense that it can be called for different argument types. - Right now this is only the case for scan operators. + This is the case for scan operators (whose argument types are `DeferredType`) and for + programs / operators taking variable-length tuples (`VarArgType`), where the concrete + tuple length is only known from the actual arguments and must drive specialization. """ + # TODO(tehrengruber): This concept does not exist elsewhere and is not properly reflected # in the type system. For now we just use `DeferredType` to communicate between # here and `type_info.type_in_program_context`. + def _contains_generic_type(type_: ts.TypeSpec) -> bool: + if isinstance(type_, (ts.DeferredType, ts.VarArgType)): + return True + if isinstance(type_, ts.TupleType): + return any(_contains_generic_type(el) for el in type_.types) + return False + return any( - isinstance(t, ts.DeferredType) + _contains_generic_type(t) for t in itertools.chain( self.program_type.definition.pos_only_args, self.program_type.definition.pos_or_kw_args.values(), diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index ac8467a5f6..1d3830b3b3 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -137,12 +137,14 @@ def tree_map_type_constructor( value: ts.CollectionTypeSpecT, elems: Iterable[ts.DataType | ts.DimensionType | ts.DeferredType], ) -> ts.CollectionTypeSpecT: + # Note: `type(value)(...)` preserves the concrete tuple subclass (e.g. `XTupleType`) so that + # element-wise-tuple semantics are not silently dropped when a tuple type is reconstructed. return ( ts.NamedCollectionType( keys=value.keys, original_python_type=value.original_python_type, types=list(elems) ) if isinstance(value, ts.NamedCollectionType) - else ts.TupleType(types=list(elems)) # type: ignore[return-value] + else type(value)(types=list(elems)) # type: ignore[return-value] ) @@ -467,10 +469,31 @@ def is_compatible_type(type_a: ts.TypeSpec, type_b: ts.TypeSpec) -> bool: is_compatible &= type_a.defined_dims == type_b.defined_dims is_compatible &= type_a.element_type == type_b.element_type elif isinstance(type_a, ts.TupleType) and isinstance(type_b, ts.TupleType): + if type(type_a) is not type(type_b): + return False if len(type_a.types) != len(type_b.types): return False for el_type_a, el_type_b in zip(type_a.types, type_b.types, strict=True): is_compatible &= is_compatible_type(el_type_a, el_type_b) + elif isinstance(type_a, ts.VarArgType) and isinstance(type_b, ts.VarArgType): + if type(type_a) is not type(type_b): + return False + is_compatible &= is_compatible_type(type_a.element_type, type_b.element_type) + elif (isinstance(type_a, ts.VarArgType) and isinstance(type_b, ts.TupleType)) or ( + isinstance(type_a, ts.TupleType) and isinstance(type_b, ts.VarArgType) + ): + # A variable-length tuple is compatible with a concrete tuple if all elements of the + # latter are compatible with the element type, e.g. when checking the concrete 'out' + # argument of a program against a field operator returning 'tuple[..., ...]'. + vararg_type, tuple_type = ( + (type_a, type_b) if isinstance(type_a, ts.VarArgType) else (type_b, type_a) + ) + assert isinstance(vararg_type, ts.VarArgType) and isinstance(tuple_type, ts.TupleType) + if isinstance(vararg_type, ts.XVarArgType) != isinstance(tuple_type, ts.XTupleType): + return False + is_compatible &= all( + is_compatible_type(vararg_type.element_type, el_type) for el_type in tuple_type.types + ) elif isinstance(type_a, ts.NamedCollectionType) and isinstance(type_b, ts.NamedCollectionType): if type_a.keys != type_b.keys: return False @@ -550,6 +573,18 @@ 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): + if type(symbol_type) is not type(to_type): + return False + return is_concretizable(symbol_type.element_type, to_type.element_type) + if isinstance(symbol_type, ts.VarArgType) and isinstance(to_type, ts.TupleType): + if isinstance(symbol_type, ts.XVarArgType) != isinstance(to_type, ts.XTupleType): + return False + 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..aeb4bdc67a 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -148,6 +148,25 @@ 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 XTupleType(TupleType): + def __str__(self) -> str: + return f"XTuple[{', '.join(map(str, self.types))}]" + + +class XVarArgType(VarArgType): + def __str__(self) -> str: + return f"XVarArgTuple[{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..c077d187be 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) @@ -198,18 +209,55 @@ def from_type_hint( ) match canonical_type: + case common.XTuple: + 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.XTupleType(types=tuple_types) + elif isinstance(args, tuple) and len(args) == 2 and args[1] is Ellipsis: + return ts.XVarArgType(element_type=from_type_hint_same_ns(args[0])) + elif args is None or (isinstance(args, tuple) and len(args) == 0): + return ts.DeferredType(constraint=ts.XTupleType) + else: + raise ValueError( + f"XTuple annotation '{type_hint}' must either " + f"be a list of concrete arguments (e.g. 'XTuple[int]'), " + f"be a variadic tuple (e.g. 'XTuple[int, ...]'), " + f"or have no arguments (e.g. 'XTuple')." + ) + 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) + 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) + 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])) + 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}'.") + 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 = ( @@ -358,6 +406,8 @@ def from_value(value: Any) -> ts.TypeSpec: # since those should be handled as general custom types. elems = [from_value(el) for el in value] assert all(isinstance(elem, ts.DataType) for elem in elems) + if isinstance(value, common.XTuple): + return ts.XTupleType(types=elems) return ts.TupleType(types=elems) elif isinstance(value, PythonNamespaceObject): return NamespaceProxy(value) diff --git a/tests/next_tests/integration_tests/cases.py b/tests/next_tests/integration_tests/cases.py index 08fb817856..1b6338a067 100644 --- a/tests/next_tests/integration_tests/cases.py +++ b/tests/next_tests/integration_tests/cases.py @@ -608,7 +608,8 @@ def _allocate_from_type( case ts.ScalarType(kind=kind): return strategy.scalar(dtype=dtype or kind.name.lower()) case ts.TupleType(types=types): - return tuple( + tuple_constructor = common.XTuple if isinstance(arg_type, ts.XTupleType) else tuple + return tuple_constructor( ( _allocate_from_type( case=case, arg_type=t, domain=domain, dtype=dtype, strategy=strategy @@ -616,6 +617,16 @@ def _allocate_from_type( for t in types ) ) + case ts.VarArgType(element_type=element_type): + tuple_constructor = common.XTuple if isinstance(arg_type, ts.XVarArgType) else tuple + return tuple_constructor( + ( + _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 +672,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_concat_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py index 577f5a520e..259bc47ab1 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,87 @@ 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_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: 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..44b6e76e73 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,210 @@ 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_var_len_tuple_comprehension_explicit_program(cartesian_case): + # A hand-written program with concrete tuple annotations wrapping a field operator + # declared with variable-length tuples. + @gtx.field_operator + def scale_tracers(tracers: tuple[cases.IField, ...], factor: int32) -> tuple[cases.IField, ...]: + return tuple(tracer * factor for tracer in tracers) + + @gtx.program + def testee( + tracers: tuple[cases.IField, cases.IField], + factor: int32, + out: tuple[cases.IField, cases.IField], + ): + scale_tracers(tracers, factor, out=out) + + 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 +def test_tuple_vararg(cartesian_case): + @gtx.field_operator + def testee( + tracers: tuple[cases.IFloatField, ...], factor: float + ) -> tuple[cases.IFloatField, cases.IFloatField]: + return tracers[0] * factor, tracers[1] * factor + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda t, f: tuple(el * f for el in t[:2]), + ) + + +@pytest.mark.uses_tuple_args +def test_elementwise_binop_fixed_xtuple(cartesian_case): + @gtx.field_operator + def testee( + a: gtx.XTuple[cases.IFloatField, cases.IFloatField], + b: gtx.XTuple[cases.IFloatField, cases.IFloatField], + ) -> gtx.XTuple[cases.IFloatField, cases.IFloatField]: + return a * b + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda a, b: tuple(x * y for x, y in zip(a, b)), + ) + + +@pytest.mark.uses_tuple_args +def test_elementwise_binop_fixed_nested_xtuple(cartesian_case): + @gtx.field_operator + def testee( + a: gtx.XTuple[gtx.XTuple[cases.IFloatField, cases.IFloatField], cases.IFloatField], + b: gtx.XTuple[cases.IFloatField, cases.IFloatField], + ) -> gtx.XTuple[gtx.XTuple[cases.IFloatField, cases.IFloatField], cases.IFloatField]: + return a * b + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda a, b: ((a[0][0] * b[0], a[0][1] * b[0]), a[1] * b[1]), + ) + + +@pytest.mark.uses_tuple_args +def test_elementwise_binop_fixed_xtuple_scalar_broadcast(cartesian_case): + @gtx.field_operator + def testee( + a: gtx.XTuple[cases.IFloatField, cases.IFloatField], factor: float + ) -> gtx.XTuple[cases.IFloatField, cases.IFloatField]: + return a * factor + a + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda a, f: tuple(el * f + el for el in a), + ) + + +@pytest.mark.uses_tuple_args +def test_elementwise_binop_var_len_xtuple_scalar_broadcast(cartesian_case): + @gtx.field_operator + def testee( + a: gtx.XTuple[cases.IFloatField, ...], factor: float + ) -> gtx.XTuple[cases.IFloatField, ...]: + return a * factor + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda a, f: tuple(el * f for el in a), + ) + + +@pytest.mark.uses_tuple_args +def test_elementwise_binop_nested_xtuple_scalar_broadcast(cartesian_case): + @gtx.field_operator + def testee( + a: gtx.XTuple[gtx.XTuple[cases.IFloatField, cases.IFloatField], cases.IFloatField], + factor: float, + ) -> gtx.XTuple[gtx.XTuple[cases.IFloatField, cases.IFloatField], cases.IFloatField]: + return a * factor + + cases.verify_with_default_data( + cartesian_case, + testee, + ref=lambda a, f: ((a[0][0] * f, a[0][1] * f), a[1] * f), + ) + + @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 +451,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_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index 407088e4f8..4d77cf47a5 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,266 @@ 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_elementwise_binop_fixed_xtuple(): + def foo( + a: gtx.XTuple[gtx.Field[[TDim], float64], gtx.Field[[TDim], float64]], + b: gtx.XTuple[gtx.Field[[TDim], float64], gtx.Field[[TDim], float64]], + ): + return a * b + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + lowered_inlined = inline_lambdas.InlineLambdas.apply(lowered) + + reference = im.make_tuple( + im.op_as_fieldop("multiplies")(im.tuple_get(0, "a"), im.tuple_get(0, "b")), + im.op_as_fieldop("multiplies")(im.tuple_get(1, "a"), im.tuple_get(1, "b")), + ) + reference_inlined = inline_lambdas.InlineLambdas.apply(reference) + + assert lowered_inlined.expr == reference_inlined + + +def test_elementwise_binop_fixed_xtuple_scalar_broadcast(): + def foo(a: gtx.XTuple[gtx.Field[[TDim], float64], gtx.Field[[TDim], float64]], factor: float64): + return a * factor + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + lowered_inlined = inline_lambdas.InlineLambdas.apply(lowered) + + reference = im.make_tuple( + im.op_as_fieldop("multiplies")(im.tuple_get(0, "a"), "factor"), + im.op_as_fieldop("multiplies")(im.tuple_get(1, "a"), "factor"), + ) + reference_inlined = inline_lambdas.InlineLambdas.apply(reference) + + assert lowered_inlined.expr == reference_inlined + + +def test_elementwise_binop_nested_xtuple_scalar_broadcast(): + def foo( + a: gtx.XTuple[ + gtx.XTuple[gtx.Field[[TDim], float64], gtx.Field[[TDim], float64]], + gtx.Field[[TDim], float64], + ], + factor: float64, + ): + return a * factor + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + lowered_inlined = inline_lambdas.InlineLambdas.apply(lowered) + + inner = im.tuple_get(0, "a") + reference = im.make_tuple( + im.make_tuple( + im.op_as_fieldop("multiplies")(im.tuple_get(0, inner), "factor"), + im.op_as_fieldop("multiplies")(im.tuple_get(1, inner), "factor"), + ), + im.op_as_fieldop("multiplies")(im.tuple_get(1, "a"), "factor"), + ) + reference_inlined = inline_lambdas.InlineLambdas.apply(reference) + + assert lowered_inlined.expr == reference_inlined + + +def test_elementwise_binop_var_len_xtuple_scalar_broadcast(): + def foo(a: gtx.XTuple[gtx.Field[[TDim], float64], ...], factor: float64): + return a * factor + + parsed = FieldOperatorParser.apply_to_function(foo) + lowered = FieldOperatorLowering.apply(parsed) + + reference = im.call( + im.call("map_tuple")( + im.lambda_("__elementwise_0")( + im.op_as_fieldop("multiplies")("__elementwise_0", "factor") + ) + ) + )("a") + + assert lowered.expr == reference + + 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/ffront_tests/test_type_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py index 22bd1a7a9e..167b847091 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 @@ -19,6 +19,7 @@ DimensionKind, Field, FieldOffset, + XTuple, astype, broadcast, errors, @@ -106,6 +107,157 @@ def nonmatching(a: Field[[X], float64], b: Field[[Y], float64]): ) +def test_binop_fixed_tuple_elementwise_product(): + def product( + f: XTuple[Field[[TDim], float64], Field[[TDim], float64]], + s: XTuple[Field[[TDim], float64], Field[[TDim], float64]], + ) -> XTuple[Field[[TDim], float64], Field[[TDim], float64]]: + return f * s + + parsed = FieldOperatorParser.apply_to_function(product) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XTupleType(types=[field_type, field_type]) + + +def test_binop_nested_fixed_tuple_elementwise_product(): + def product( + f: XTuple[XTuple[Field[[TDim], float64], Field[[TDim], float64]], Field[[TDim], float64]], + s: XTuple[XTuple[Field[[TDim], float64], Field[[TDim], float64]], Field[[TDim], float64]], + ) -> XTuple[XTuple[Field[[TDim], float64], Field[[TDim], float64]], Field[[TDim], float64]]: + return f * s + + parsed = FieldOperatorParser.apply_to_function(product) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XTupleType( + types=[ts.XTupleType(types=[field_type, field_type]), field_type] + ) + + +def test_binop_nested_same_outer_fixed_tuple_elementwise_product(): + def product( + f: XTuple[XTuple[Field[[TDim], float64], Field[[TDim], float64]], Field[[TDim], float64]], + s: XTuple[Field[[TDim], float64], Field[[TDim], float64]], + ) -> XTuple[XTuple[Field[[TDim], float64], Field[[TDim], float64]], Field[[TDim], float64]]: + return f * s + + parsed = FieldOperatorParser.apply_to_function(product) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XTupleType( + types=[ts.XTupleType(types=[field_type, field_type]), field_type] + ) + + +def test_binop_fixed_tuple_scalar_elementwise_product(): + def product( + f: XTuple[Field[[TDim], float64], Field[[TDim], float64]], s: float64 + ) -> XTuple[Field[[TDim], float64], Field[[TDim], float64]]: + return f * s + + parsed = FieldOperatorParser.apply_to_function(product) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XTupleType(types=[field_type, field_type]) + + +def test_binop_fixed_nested_tuple_scalar_elementwise_product(): + def product( + f: XTuple[XTuple[Field[[TDim], float64], Field[[TDim], float64]], Field[[TDim], float64]], + s: float64, + ) -> XTuple[XTuple[Field[[TDim], float64], Field[[TDim], float64]], Field[[TDim], float64]]: + return f * s + + parsed = FieldOperatorParser.apply_to_function(product) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XTupleType( + types=[ts.XTupleType(types=[field_type, field_type]), field_type] + ) + + +def test_binop_fixed_tuple_field_elementwise_product(): + def product( + f: XTuple[Field[[TDim], float64], Field[[TDim], float64]], s: Field[[TDim], float64] + ) -> XTuple[Field[[TDim], float64], Field[[TDim], float64]]: + return f * s + + parsed = FieldOperatorParser.apply_to_function(product) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XTupleType(types=[field_type, field_type]) + + +def test_binop_var_len_tuple_scalar_elementwise_product(): + def product( + f: XTuple[Field[[TDim], float64], ...], s: float64 + ) -> XTuple[Field[[TDim], float64], ...]: + return f * s + + parsed = FieldOperatorParser.apply_to_function(product) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XVarArgType(element_type=field_type) + + +def test_binop_var_len_tuple_elementwise_product_unsupported(): + def product( + f: XTuple[Field[[TDim], float64], ...], s: XTuple[Field[[TDim], float64], ...] + ) -> XTuple[Field[[TDim], float64], ...]: + return f * s + + with pytest.raises(errors.DSLError, match="two variable-length tuples"): + _ = FieldOperatorParser.apply_to_function(product) + + +def test_binop_tuple_structure_mismatch(): + def product( + f: XTuple[Field[[TDim], float64], Field[[TDim], float64]], + s: XTuple[Field[[TDim], float64]], + ): + return f * s + + with pytest.raises(errors.DSLError, match="same structure"): + _ = FieldOperatorParser.apply_to_function(product) + + +def test_binop_regular_tuple_unsupported(): + def product(f: tuple[Field[[TDim], float64], Field[[TDim], float64]], s: float64): + return f * s + + with pytest.raises(errors.DSLError, match="Unsupported operand type"): + _ = FieldOperatorParser.apply_to_function(product) + + +def test_vararg_subscript_requires_literal_index(): + def foo(t: tuple[Field[[TDim], float64], ...], i: int32): + return t[i] + + with pytest.raises(errors.DSLError, match="indexed with literal integers"): + _ = FieldOperatorParser.apply_to_function(foo) + + +def test_tuple_comprehension_over_xtuple_preserves_xtuple(): + def foo(f: XTuple[Field[[TDim], float64], Field[[TDim], float64]]): + return tuple(2.0 * el for el in f) + + parsed = FieldOperatorParser.apply_to_function(foo) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XTupleType(types=[field_type, field_type]) + + +def test_tuple_comprehension_over_xvararg_preserves_xtuple(): + def foo(f: XTuple[Field[[TDim], float64], ...]): + return tuple(2.0 * el for el in f) + + parsed = FieldOperatorParser.apply_to_function(foo) + + field_type = ts.FieldType(dims=[TDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) + assert parsed.body.stmts[0].value.type == ts.XVarArgType(element_type=field_type) + + def test_bitop_float(): def float_bitop(a: Field[[TDim], float], b: Field[[TDim], float]): return a & b 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..2acb6af1f1 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,49 @@ 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_xtuple_is_not_compatible_with_regular_tuple(): + scalar_type = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) + + assert not type_info.is_compatible_type( + ts.XTupleType(types=[scalar_type]), ts.TupleType(types=[scalar_type]) + ) + assert not type_info.is_compatible_type( + ts.TupleType(types=[scalar_type]), ts.XTupleType(types=[scalar_type]) + ) + assert not type_info.is_compatible_type( + ts.XVarArgType(element_type=scalar_type), ts.VarArgType(element_type=scalar_type) + ) + assert not type_info.is_compatible_type( + ts.VarArgType(element_type=scalar_type), ts.XVarArgType(element_type=scalar_type) + ) + assert not type_info.is_concretizable( + ts.XVarArgType(element_type=scalar_type), ts.TupleType(types=[scalar_type]) + ) + assert not type_info.is_concretizable( + ts.VarArgType(element_type=scalar_type), ts.XTupleType(types=[scalar_type]) + ) + assert type_info.is_concretizable( + ts.XVarArgType(element_type=scalar_type), ts.XTupleType(types=[scalar_type]) + ) + + +def test_vararg_is_compatible_with_uniform_tuple(): + scalar_type = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) + other_type = ts.ScalarType(kind=ts.ScalarKind.INT32) + vararg_type = ts.VarArgType(element_type=scalar_type) + + assert type_info.is_compatible_type(vararg_type, ts.TupleType(types=[scalar_type, scalar_type])) + assert type_info.is_compatible_type(ts.TupleType(types=[scalar_type, scalar_type]), vararg_type) + assert type_info.is_compatible_type(vararg_type, ts.TupleType(types=[])) + assert not type_info.is_compatible_type( + vararg_type, ts.TupleType(types=[scalar_type, other_type]) + ) + assert not type_info.is_compatible_type(vararg_type, ts.XTupleType(types=[scalar_type])) + assert not type_info.is_compatible_type( + ts.XVarArgType(element_type=scalar_type), ts.TupleType(types=[scalar_type]) + ) + assert type_info.is_compatible_type( + ts.XVarArgType(element_type=scalar_type), ts.XTupleType(types=[scalar_type]) + ) 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..c8c441243d 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 @@ -93,6 +93,20 @@ def _make_type_string_for_container(cls: type) -> str: ] ), ), + ( + gtx.XTuple[bool, typing.Tuple[int, float]], + ts.XTupleType( + types=[ + ts.ScalarType(kind=ts.ScalarKind.BOOL), + ts.TupleType( + types=[ + ts.ScalarType(kind=ts.ScalarKind.INT64), + ts.ScalarType(kind=ts.ScalarKind.FLOAT64), + ] + ), + ] + ), + ), ( gtx.Field[[IDim], float], ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)), @@ -216,15 +230,20 @@ 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) + ) + assert type_translation.from_type_hint(gtx.XTuple[int, ...]) == ts.XVarArgType( + element_type=ts.ScalarType(kind=ts.ScalarKind.INT64) + ) # Fields with pytest.raises(ValueError, match="Field type requires two arguments"):