From c247d4110e4f15b4f76327fa11b138795a75325d Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 3 Oct 2024 18:14:39 +0200 Subject: [PATCH 01/48] Add prototype support for for loops over data dims. --- .../cartesian/frontend/gtscript_frontend.py | 67 +++++++++++++++++++ src/gt4py/cartesian/gtscript.py | 1 + 2 files changed, 68 insertions(+) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index e2aa98f3cf..c1c56f1e04 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -613,6 +613,70 @@ def visit_If(self, node: ast.If): return node if node else None +class DataDimLoopIndexReplacer(ast.NodeTransformer): + def __init__(self, name: str, value: int) -> None: + self.name = name + self.value = value + + def visit_Name(self, node: ast.Name) -> Union[ast.Constant, ast.Name]: + if node.id == self.name: + return ast.Constant(self.value, ctx=node.ctx) + else: + return node + + +class DataDimLoopUnroller(ast.NodeTransformer): + @classmethod + def apply(cls, func_node: ast.FunctionDef, context: dict): + unroller = cls(context) + unroller(func_node) + + def __init__(self, context): + self.context = context + self.prefix = "" + + def __call__(self, func_node: ast.FunctionDef): + self.visit(func_node) + + def visit_For(self, node: ast.For) -> Union[ast.For, list[ast.AST]]: + super().generic_visit(node) + + if ( + isinstance(node.iter, ast.Call) + and isinstance(node.iter.func, ast.Name) + and node.iter.func.id == "range" + ): + range_args = node.iter.args + assert all(isinstance(arg, ast.Constant) for arg in range_args) + assert 1 <= len(range_args) <= 3 + if len(range_args) == 1: + start = 0 + stop = range_args[0].value + step = 1 + elif len(range_args) == 2: + start = range_args[0].value + stop = range_args[1].value + step = 1 + else: + start = range_args[0].value + stop = range_args[1].value + step = range_args[2].value + + assert isinstance(node.target, ast.Name) + index_name = node.target.id + + new_body = [] + for i in range(start, stop, step): + body = copy.deepcopy(node.body) + transformer = DataDimLoopIndexReplacer(index_name, i) + new_body_item = [transformer.visit(stmt) for stmt in body] + new_body += new_body_item + + return new_body + else: + return node + + def _make_temp_decls( descriptors: Dict[str, gtscript._FieldDescriptor], ) -> Dict[str, nodes.FieldDecl]: @@ -2055,6 +2119,9 @@ def run(self, backend_name: str): ValueInliner.apply(main_func_node, context=local_context) + # unroll loops over data dimensions + DataDimLoopUnroller.apply(main_func_node, context=local_context) + # Inline function calls CallInliner.apply(main_func_node, context=local_context) diff --git a/src/gt4py/cartesian/gtscript.py b/src/gt4py/cartesian/gtscript.py index 643ecba010..b823f960cd 100644 --- a/src/gt4py/cartesian/gtscript.py +++ b/src/gt4py/cartesian/gtscript.py @@ -83,6 +83,7 @@ "__externals__", "__INLINED", "compile_assert", + "range", *MATH_BUILTINS, } From cc2def32fd2dccdf9fd179ab4e58238a51fe74bf Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 4 Oct 2024 12:09:13 +0200 Subject: [PATCH 02/48] Add support for lists and tuples. --- src/gt4py/cartesian/frontend/gtscript_frontend.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index c1c56f1e04..0a38eee176 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -641,6 +641,7 @@ def __call__(self, func_node: ast.FunctionDef): def visit_For(self, node: ast.For) -> Union[ast.For, list[ast.AST]]: super().generic_visit(node) + index_values: Optional[Union[list, range]] = None if ( isinstance(node.iter, ast.Call) and isinstance(node.iter.func, ast.Name) @@ -661,14 +662,20 @@ def visit_For(self, node: ast.For) -> Union[ast.For, list[ast.AST]]: start = range_args[0].value stop = range_args[1].value step = range_args[2].value + index_values = range(start, stop, step) + elif isinstance(node.iter, (ast.List, ast.Tuple)): + index_value_nodes = node.iter.elts + assert all(isinstance(node, ast.Constant) for node in index_value_nodes) + index_values = [node.value for node in index_value_nodes] + if index_values is not None: assert isinstance(node.target, ast.Name) index_name = node.target.id new_body = [] - for i in range(start, stop, step): + for index_value in index_values: body = copy.deepcopy(node.body) - transformer = DataDimLoopIndexReplacer(index_name, i) + transformer = DataDimLoopIndexReplacer(index_name, index_value) new_body_item = [transformer.visit(stmt) for stmt in body] new_body += new_body_item From e79917af097efacc32b25339b03f25fc83abf464 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 24 Oct 2024 11:29:39 +0200 Subject: [PATCH 03/48] Cosmetics. --- .../cartesian/frontend/gtscript_frontend.py | 19 +++++-------------- 1 file changed, 5 insertions(+), 14 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 0a38eee176..be90a05414 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -647,26 +647,17 @@ def visit_For(self, node: ast.For) -> Union[ast.For, list[ast.AST]]: and isinstance(node.iter.func, ast.Name) and node.iter.func.id == "range" ): - range_args = node.iter.args - assert all(isinstance(arg, ast.Constant) for arg in range_args) + range_args = [eval(ast.unparse(arg), self.context) for arg in node.iter.args] assert 1 <= len(range_args) <= 3 if len(range_args) == 1: - start = 0 - stop = range_args[0].value - step = 1 + start, stop, step = 0, *range_args, 1 elif len(range_args) == 2: - start = range_args[0].value - stop = range_args[1].value - step = 1 + start, stop, step = *range_args, 1 else: - start = range_args[0].value - stop = range_args[1].value - step = range_args[2].value + start, stop, step = range_args index_values = range(start, stop, step) elif isinstance(node.iter, (ast.List, ast.Tuple)): - index_value_nodes = node.iter.elts - assert all(isinstance(node, ast.Constant) for node in index_value_nodes) - index_values = [node.value for node in index_value_nodes] + index_values = [eval(ast.unparse(elt), self.context) for elt in node.iter.elts] if index_values is not None: assert isinstance(node.target, ast.Name) From 0d321ecac0b4de4ac2fc27b7abf7b6e25565e73a Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 24 Oct 2024 11:30:49 +0200 Subject: [PATCH 04/48] Allow field data indices. --- src/gt4py/cartesian/gtc/numpy/npir.py | 8 -------- src/gt4py/cartesian/gtc/numpy/npir_codegen.py | 3 ++- 2 files changed, 2 insertions(+), 9 deletions(-) diff --git a/src/gt4py/cartesian/gtc/numpy/npir.py b/src/gt4py/cartesian/gtc/numpy/npir.py index 6532a2789e..36a17b7301 100644 --- a/src/gt4py/cartesian/gtc/numpy/npir.py +++ b/src/gt4py/cartesian/gtc/numpy/npir.py @@ -126,14 +126,6 @@ class FieldSlice(VectorLValue): data_index: List[Expr] = eve.field(default_factory=list) kind: common.ExprKind = common.ExprKind.FIELD - @datamodels.validator("data_index") - def data_indices_are_scalar( - self, attribute: datamodels.Attribute, data_index: List[Expr] - ) -> None: - for index in data_index: - if index.kind != common.ExprKind.SCALAR: - raise ValueError("Data indices must be scalars") - class ParamAccess(Expr): name: eve.Coerced[eve.SymbolRef] diff --git a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py index e1a9f8e8bb..40bd75e534 100644 --- a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py +++ b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py @@ -187,7 +187,8 @@ def visit_FieldSlice(self, node: npir.FieldSlice, **kwargs: Any) -> Union[str, C ) args = _make_slice_access(offsets, kwargs["is_serial"], kwargs.get("horizontal_mask")) - data_index = self.visit(node.data_index, inside_slice=True, **kwargs) + kwargs["inside_slice"] = True + data_index = self.visit(node.data_index, **kwargs) access_slice = ", ".join(args + list(data_index)) From 0856861f3d5e6732fe95affd3c68c8683c223c72 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 25 Oct 2024 10:14:08 +0200 Subject: [PATCH 05/48] Call CallInliner before DataDimLoopUnroller. --- src/gt4py/cartesian/frontend/gtscript_frontend.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index be90a05414..5ab00dea4e 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -2117,12 +2117,12 @@ def run(self, backend_name: str): ValueInliner.apply(main_func_node, context=local_context) - # unroll loops over data dimensions - DataDimLoopUnroller.apply(main_func_node, context=local_context) - # Inline function calls CallInliner.apply(main_func_node, context=local_context) + # unroll loops over data dimensions + DataDimLoopUnroller.apply(main_func_node, context=local_context) + # Evaluate and inline compile-time conditionals CompiledIfInliner.apply(main_func_node, context=local_context, stencil_name=self.main_name) From 4847fb4e3defbbe71cf9ebce6ef6ae69d0b93710 Mon Sep 17 00:00:00 2001 From: Florian Deconinck Date: Mon, 19 Aug 2024 15:06:41 -0400 Subject: [PATCH 06/48] Casting to INT. Add `v_in_int = int(v_in_float)` op as a base unitary op. Deactivate upcaster for thi cast call. --- src/gt4py/cartesian/frontend/defir_to_gtir.py | 1 + src/gt4py/cartesian/frontend/gtscript_frontend.py | 1 + src/gt4py/cartesian/frontend/nodes.py | 3 +++ src/gt4py/cartesian/gtc/common.py | 6 ++++++ src/gt4py/cartesian/gtc/cuir/cuir_codegen.py | 1 + src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py | 1 + src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py | 1 + src/gt4py/cartesian/gtc/passes/gtir_upcaster.py | 5 ++++- src/gt4py/cartesian/gtc/ufuncs.py | 1 + src/gt4py/cartesian/gtscript.py | 1 + 10 files changed, 20 insertions(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/frontend/defir_to_gtir.py b/src/gt4py/cartesian/frontend/defir_to_gtir.py index 5d38e077fb..dabc74d246 100644 --- a/src/gt4py/cartesian/frontend/defir_to_gtir.py +++ b/src/gt4py/cartesian/frontend/defir_to_gtir.py @@ -331,6 +331,7 @@ class DefIRToGTIR(IRNodeVisitor): NativeFunction.FLOOR: common.NativeFunction.FLOOR, NativeFunction.CEIL: common.NativeFunction.CEIL, NativeFunction.TRUNC: common.NativeFunction.TRUNC, + NativeFunction.INT: common.NativeFunction.INT, } GT4PY_BUILTIN_TO_GTIR = { diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 5ab00dea4e..17a84b4032 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -829,6 +829,7 @@ def __init__( "floor": nodes.NativeFunction.FLOOR, "ceil": nodes.NativeFunction.CEIL, "trunc": nodes.NativeFunction.TRUNC, + "int": nodes.NativeFunction.INT, } def __call__(self, ast_root: ast.AST): diff --git a/src/gt4py/cartesian/frontend/nodes.py b/src/gt4py/cartesian/frontend/nodes.py index f84577e7b5..82ddae5f5e 100644 --- a/src/gt4py/cartesian/frontend/nodes.py +++ b/src/gt4py/cartesian/frontend/nodes.py @@ -411,6 +411,8 @@ class NativeFunction(enum.Enum): CEIL = enum.auto() TRUNC = enum.auto() + INT = enum.auto() + @property def arity(self): return type(self).IR_OP_TO_NUM_ARGS[self] @@ -445,6 +447,7 @@ def arity(self): NativeFunction.FLOOR: 1, NativeFunction.CEIL: 1, NativeFunction.TRUNC: 1, + NativeFunction.INT: 1, } diff --git a/src/gt4py/cartesian/gtc/common.py b/src/gt4py/cartesian/gtc/common.py index bfe434e7f3..bac183bfb1 100644 --- a/src/gt4py/cartesian/gtc/common.py +++ b/src/gt4py/cartesian/gtc/common.py @@ -178,6 +178,8 @@ class NativeFunction(eve.StrEnum): CEIL = "ceil" TRUNC = "trunc" + INT = "int" + IR_OP_TO_NUM_ARGS: ClassVar[Dict[NativeFunction, int]] @property @@ -217,6 +219,7 @@ def arity(self) -> int: NativeFunction.FLOOR: 1, NativeFunction.CEIL: 1, NativeFunction.TRUNC: 1, + NativeFunction.INT: 1, }.items() } @@ -551,6 +554,8 @@ def native_func_call_dtype_propagation(*, strict: bool = True) -> datamodels.Roo def _impl(cls: Type[NativeFuncCall], instance: NativeFuncCall) -> None: if instance.func in (NativeFunction.ISFINITE, NativeFunction.ISINF, NativeFunction.ISNAN): instance.dtype = DataType.BOOL # type: ignore[attr-defined] + elif instance.func in (NativeFunction.INT): + instance.dtype = DataType.INT32 else: # assumes all NativeFunction args have a common dtype common_dtype = verify_and_get_common_dtype(cls, instance.args, strict=strict) @@ -887,6 +892,7 @@ def data_type_to_typestr(dtype: DataType) -> str: NativeFunction.FLOOR: "floor", NativeFunction.CEIL: "ceil", NativeFunction.TRUNC: "trunc", + NativeFunction.INT: "int", }, } diff --git a/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py b/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py index 76f076874a..1ba3ce3f9e 100644 --- a/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py +++ b/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py @@ -169,6 +169,7 @@ def visit_Literal( NativeFunction.FLOOR: "std::floor", NativeFunction.CEIL: "std::ceil", NativeFunction.TRUNC: "std::trunc", + NativeFunction.INT: "int", } def visit_NativeFunction(self, func: NativeFunction, **kwargs: Any) -> str: diff --git a/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py b/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py index 696dc27387..fcbba6c32c 100644 --- a/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py +++ b/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py @@ -167,6 +167,7 @@ def visit_NativeFunction(self, func: common.NativeFunction, **kwargs: Any) -> st common.NativeFunction.FLOOR: "dace.math.ifloor", common.NativeFunction.CEIL: "ceil", common.NativeFunction.TRUNC: "trunc", + common.NativeFunction.INT: "int", }[func] except KeyError as error: raise NotImplementedError("Not implemented NativeFunction encountered.") from error diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py index 3105f4a8cb..d513e29838 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py @@ -176,6 +176,7 @@ def visit_NativeFunction(self, func: NativeFunction, **kwargs: Any) -> str: NativeFunction.FLOOR: "std::floor", NativeFunction.CEIL: "std::ceil", NativeFunction.TRUNC: "std::trunc", + NativeFunction.INT: "int", }[func] except KeyError as error: raise NotImplementedError( diff --git a/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py b/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py index 41fa127d6d..6cf3e567cd 100644 --- a/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py +++ b/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py @@ -13,7 +13,7 @@ from gt4py import eve from gt4py.cartesian.gtc import gtir -from gt4py.cartesian.gtc.common import DataType, op_to_ufunc, typestr_to_data_type +from gt4py.cartesian.gtc.common import DataType, NativeFunction, op_to_ufunc, typestr_to_data_type from gt4py.cartesian.gtc.gtir import Expr from gt4py.eve import datamodels @@ -104,6 +104,9 @@ def visit_TernaryOp(self, node: gtir.TernaryOp, **kwargs: Any) -> gtir.TernaryOp ) def visit_NativeFuncCall(self, node: gtir.NativeFuncCall, **kwargs: Any) -> gtir.NativeFuncCall: + # Skip upcasting for cast to int + if node.func == NativeFunction.INT: + return node upcasting_rule = functools.partial( _numpy_ufunc_upcasting_rule, ufunc=op_to_ufunc(node.func) ) diff --git a/src/gt4py/cartesian/gtc/ufuncs.py b/src/gt4py/cartesian/gtc/ufuncs.py index 88c7534602..74d49394ab 100644 --- a/src/gt4py/cartesian/gtc/ufuncs.py +++ b/src/gt4py/cartesian/gtc/ufuncs.py @@ -63,3 +63,4 @@ floor: np.ufunc = np.floor ceil: np.ufunc = np.ceil trunc: np.ufunc = np.trunc +int: np.ufunc = np.int32 # noqa: A001 [builtin-variable-shadowing] diff --git a/src/gt4py/cartesian/gtscript.py b/src/gt4py/cartesian/gtscript.py index b823f960cd..36e64d778d 100644 --- a/src/gt4py/cartesian/gtscript.py +++ b/src/gt4py/cartesian/gtscript.py @@ -59,6 +59,7 @@ "floor", "ceil", "trunc", + "int", } builtins = { From 8f5fe956e466ec98a1ce4f4e60abe7856dd130db Mon Sep 17 00:00:00 2001 From: Florian Deconinck Date: Thu, 10 Oct 2024 14:01:40 -0400 Subject: [PATCH 07/48] Lint --- src/gt4py/cartesian/gtc/common.py | 2 +- src/gt4py/cartesian/gtc/ufuncs.py | 4 +++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/src/gt4py/cartesian/gtc/common.py b/src/gt4py/cartesian/gtc/common.py index bac183bfb1..a62c4c2bd0 100644 --- a/src/gt4py/cartesian/gtc/common.py +++ b/src/gt4py/cartesian/gtc/common.py @@ -555,7 +555,7 @@ def _impl(cls: Type[NativeFuncCall], instance: NativeFuncCall) -> None: if instance.func in (NativeFunction.ISFINITE, NativeFunction.ISINF, NativeFunction.ISNAN): instance.dtype = DataType.BOOL # type: ignore[attr-defined] elif instance.func in (NativeFunction.INT): - instance.dtype = DataType.INT32 + instance.dtype = DataType.INT32 # type: ignore[attr-defined] else: # assumes all NativeFunction args have a common dtype common_dtype = verify_and_get_common_dtype(cls, instance.args, strict=strict) diff --git a/src/gt4py/cartesian/gtc/ufuncs.py b/src/gt4py/cartesian/gtc/ufuncs.py index 74d49394ab..be5f78fdcd 100644 --- a/src/gt4py/cartesian/gtc/ufuncs.py +++ b/src/gt4py/cartesian/gtc/ufuncs.py @@ -6,6 +6,8 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +from typing import Type + import numpy as np @@ -63,4 +65,4 @@ floor: np.ufunc = np.floor ceil: np.ufunc = np.ceil trunc: np.ufunc = np.trunc -int: np.ufunc = np.int32 # noqa: A001 [builtin-variable-shadowing] +int: Type[np.signedinteger] = np.int32 # noqa: A001 [builtin-variable-shadowing] From a70a46c940ab550964215c81577606f9879e8982 Mon Sep 17 00:00:00 2001 From: Florian Deconinck Date: Thu, 24 Oct 2024 16:34:56 -0400 Subject: [PATCH 08/48] Native function: f32, f64, round --- src/gt4py/cartesian/frontend/defir_to_gtir.py | 3 +++ .../cartesian/frontend/gtscript_frontend.py | 3 +++ src/gt4py/cartesian/frontend/nodes.py | 9 ++++++++- src/gt4py/cartesian/gtc/common.py | 16 +++++++++++++++- src/gt4py/cartesian/gtc/cuir/cuir_codegen.py | 3 +++ .../gtc/dace/expansion/tasklet_codegen.py | 3 +++ src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py | 3 +++ src/gt4py/cartesian/gtc/passes/gtir_upcaster.py | 2 +- src/gt4py/cartesian/gtc/ufuncs.py | 3 +++ src/gt4py/cartesian/gtscript.py | 11 ++++++++++- 10 files changed, 52 insertions(+), 4 deletions(-) diff --git a/src/gt4py/cartesian/frontend/defir_to_gtir.py b/src/gt4py/cartesian/frontend/defir_to_gtir.py index dabc74d246..ede29efc6f 100644 --- a/src/gt4py/cartesian/frontend/defir_to_gtir.py +++ b/src/gt4py/cartesian/frontend/defir_to_gtir.py @@ -331,7 +331,10 @@ class DefIRToGTIR(IRNodeVisitor): NativeFunction.FLOOR: common.NativeFunction.FLOOR, NativeFunction.CEIL: common.NativeFunction.CEIL, NativeFunction.TRUNC: common.NativeFunction.TRUNC, + NativeFunction.ROUND: common.NativeFunction.ROUND, NativeFunction.INT: common.NativeFunction.INT, + NativeFunction.F64: common.NativeFunction.F64, + NativeFunction.F32: common.NativeFunction.F32, } GT4PY_BUILTIN_TO_GTIR = { diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 17a84b4032..88cda797a7 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -829,7 +829,10 @@ def __init__( "floor": nodes.NativeFunction.FLOOR, "ceil": nodes.NativeFunction.CEIL, "trunc": nodes.NativeFunction.TRUNC, + "round": nodes.NativeFunction.ROUND, "int": nodes.NativeFunction.INT, + "f32": nodes.NativeFunction.F32, + "f64": nodes.NativeFunction.F64, } def __call__(self, ast_root: ast.AST): diff --git a/src/gt4py/cartesian/frontend/nodes.py b/src/gt4py/cartesian/frontend/nodes.py index 82ddae5f5e..01e7dd57f4 100644 --- a/src/gt4py/cartesian/frontend/nodes.py +++ b/src/gt4py/cartesian/frontend/nodes.py @@ -40,7 +40,8 @@ NativeFunction enumeration (:class:`NativeFunction`) Native function identifier [`ABS`, `MAX`, `MIN, `MOD`, `SIN`, `COS`, `TAN`, `ARCSIN`, `ARCCOS`, `ARCTAN`, - `SQRT`, `EXP`, `LOG`, `LOG10`, `ISFINITE`, `ISINF`, `ISNAN`, `FLOOR`, `CEIL`, `TRUNC`] + `SQRT`, `EXP`, `LOG`, `LOG10`, `ISFINITE`, `ISINF`, `ISNAN`, `FLOOR`, `CEIL`, `TRUNC` + `ROUND`, `INT`, `F32`, `F64`] LevelMarker enumeration (:class:`LevelMarker`) Special axis levels @@ -410,8 +411,11 @@ class NativeFunction(enum.Enum): FLOOR = enum.auto() CEIL = enum.auto() TRUNC = enum.auto() + ROUND = enum.auto() INT = enum.auto() + F32 = enum.auto() + F64 = enum.auto() @property def arity(self): @@ -447,7 +451,10 @@ def arity(self): NativeFunction.FLOOR: 1, NativeFunction.CEIL: 1, NativeFunction.TRUNC: 1, + NativeFunction.ROUND: 1, NativeFunction.INT: 1, + NativeFunction.F32: 1, + NativeFunction.F64: 1, } diff --git a/src/gt4py/cartesian/gtc/common.py b/src/gt4py/cartesian/gtc/common.py index a62c4c2bd0..e78614c5b0 100644 --- a/src/gt4py/cartesian/gtc/common.py +++ b/src/gt4py/cartesian/gtc/common.py @@ -177,8 +177,11 @@ class NativeFunction(eve.StrEnum): FLOOR = "floor" CEIL = "ceil" TRUNC = "trunc" + ROUND = "round" INT = "int" + F32 = "f32" + F64 = "f64" IR_OP_TO_NUM_ARGS: ClassVar[Dict[NativeFunction, int]] @@ -219,7 +222,10 @@ def arity(self) -> int: NativeFunction.FLOOR: 1, NativeFunction.CEIL: 1, NativeFunction.TRUNC: 1, + NativeFunction.ROUND: 1, NativeFunction.INT: 1, + NativeFunction.F32: 1, + NativeFunction.F64: 1, }.items() } @@ -605,7 +611,12 @@ def visit_Node( self.generic_visit(node, loop_order=loop_order, **kwargs) def visit_AssignStmt( - self, node: AssignStmt, *, loop_order: LoopOrder, symtable: Dict[str, Any], **kwargs: Any + self, + node: AssignStmt, + *, + loop_order: LoopOrder, + symtable: Dict[str, Any], + **kwargs: Any, ) -> None: decl = symtable.get(node.left.name, None) if decl is None: @@ -892,7 +903,10 @@ def data_type_to_typestr(dtype: DataType) -> str: NativeFunction.FLOOR: "floor", NativeFunction.CEIL: "ceil", NativeFunction.TRUNC: "trunc", + NativeFunction.TRUNC: "round", NativeFunction.INT: "int", + NativeFunction.F32: "f32", + NativeFunction.F64: "f64", }, } diff --git a/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py b/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py index 1ba3ce3f9e..ce2775384c 100644 --- a/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py +++ b/src/gt4py/cartesian/gtc/cuir/cuir_codegen.py @@ -169,7 +169,10 @@ def visit_Literal( NativeFunction.FLOOR: "std::floor", NativeFunction.CEIL: "std::ceil", NativeFunction.TRUNC: "std::trunc", + NativeFunction.ROUND: "std::round", NativeFunction.INT: "int", + NativeFunction.F32: "float", + NativeFunction.F64: "double", } def visit_NativeFunction(self, func: NativeFunction, **kwargs: Any) -> str: diff --git a/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py b/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py index fcbba6c32c..6cd2ea2044 100644 --- a/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py +++ b/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py @@ -167,7 +167,10 @@ def visit_NativeFunction(self, func: common.NativeFunction, **kwargs: Any) -> st common.NativeFunction.FLOOR: "dace.math.ifloor", common.NativeFunction.CEIL: "ceil", common.NativeFunction.TRUNC: "trunc", + common.NativeFunction.ROUND: "round", common.NativeFunction.INT: "int", + common.NativeFunction.F32: "dace.float32", + common.NativeFunction.F64: "dace.float64", }[func] except KeyError as error: raise NotImplementedError("Not implemented NativeFunction encountered.") from error diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py index d513e29838..d9790249d9 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py @@ -176,7 +176,10 @@ def visit_NativeFunction(self, func: NativeFunction, **kwargs: Any) -> str: NativeFunction.FLOOR: "std::floor", NativeFunction.CEIL: "std::ceil", NativeFunction.TRUNC: "std::trunc", + NativeFunction.ROUND: "std::round", NativeFunction.INT: "int", + NativeFunction.F32: "float", + NativeFunction.F64: "double", }[func] except KeyError as error: raise NotImplementedError( diff --git a/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py b/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py index 6cf3e567cd..24a4287db8 100644 --- a/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py +++ b/src/gt4py/cartesian/gtc/passes/gtir_upcaster.py @@ -105,7 +105,7 @@ def visit_TernaryOp(self, node: gtir.TernaryOp, **kwargs: Any) -> gtir.TernaryOp def visit_NativeFuncCall(self, node: gtir.NativeFuncCall, **kwargs: Any) -> gtir.NativeFuncCall: # Skip upcasting for cast to int - if node.func == NativeFunction.INT: + if node.func in [NativeFunction.INT, NativeFunction.F32, NativeFunction.F64]: return node upcasting_rule = functools.partial( _numpy_ufunc_upcasting_rule, ufunc=op_to_ufunc(node.func) diff --git a/src/gt4py/cartesian/gtc/ufuncs.py b/src/gt4py/cartesian/gtc/ufuncs.py index be5f78fdcd..e61f307512 100644 --- a/src/gt4py/cartesian/gtc/ufuncs.py +++ b/src/gt4py/cartesian/gtc/ufuncs.py @@ -65,4 +65,7 @@ floor: np.ufunc = np.floor ceil: np.ufunc = np.ceil trunc: np.ufunc = np.trunc +round: np.ufunc = np.round int: Type[np.signedinteger] = np.int32 # noqa: A001 [builtin-variable-shadowing] +f32: Type[np.floating] = np.float32 # type : ignore +f64: Type[np.floating] = np.float64 diff --git a/src/gt4py/cartesian/gtscript.py b/src/gt4py/cartesian/gtscript.py index 36e64d778d..30700e6246 100644 --- a/src/gt4py/cartesian/gtscript.py +++ b/src/gt4py/cartesian/gtscript.py @@ -59,7 +59,10 @@ "floor", "ceil", "trunc", + "round", "int", + "f32", + "f64", } builtins = { @@ -276,7 +279,13 @@ def stencil( # Setup build_info timings if build_info is not None: - time_keys = ("parse_time", "module_time", "codegen_time", "build_time", "load_time") + time_keys = ( + "parse_time", + "module_time", + "codegen_time", + "build_time", + "load_time", + ) build_info.update({time_key: 0.0 for time_key in time_keys}) build_options = gt_definitions.BuildOptions( From 1d7cbadc07d316125d6f8eb8c7c268a30df231b6 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 25 Oct 2024 21:48:19 +0200 Subject: [PATCH 09/48] Enable for-loops in functions. --- src/gt4py/cartesian/gtscript.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/gtscript.py b/src/gt4py/cartesian/gtscript.py index 30700e6246..bbe88e2811 100644 --- a/src/gt4py/cartesian/gtscript.py +++ b/src/gt4py/cartesian/gtscript.py @@ -91,7 +91,7 @@ *MATH_BUILTINS, } -IGNORE_WHEN_INLINING = {*MATH_BUILTINS, "compile_assert"} +IGNORE_WHEN_INLINING = {*MATH_BUILTINS, "compile_assert", "range"} __all__ = [*list(builtins), "function", "stencil", "lazy_stencil"] From 511222a23725780fffa5d7905ae54f4421e9f0bd Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 25 Oct 2024 21:48:45 +0200 Subject: [PATCH 10/48] Enable functions with no return statements. --- .../cartesian/frontend/gtscript_frontend.py | 52 ++++++++++++------- 1 file changed, 32 insertions(+), 20 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 88cda797a7..bed8fb3418 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -309,11 +309,15 @@ def visit_FunctionDef(self, node: ast.FunctionDef): class ReturnReplacer(gt_utils.meta.ASTTransformPass): @classmethod - def apply(cls, ast_object: ast.AST, target_node: ast.AST) -> None: + def apply(cls, ast_object: ast.AST, target_node: Optional[ast.AST]) -> None: """Ensure that there is only a single return statement (can still return a tuple).""" ret_count = sum(isinstance(node, ast.Return) for node in ast.walk(ast_object)) - if ret_count != 1: - raise GTScriptSyntaxError("GTScript Functions should have a single return statement") + if ret_count > 1: + raise GTScriptSyntaxError("GTScript Functions cannot have multiple return statements") + elif ret_count == 0 and target_node is not None: + raise GTScriptSyntaxError( + "Attempting to assign the return value of a GTScript function that does not return anything." + ) cls().visit(ast_object, target_node=target_node) @staticmethod @@ -493,26 +497,25 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex # Replace returns by assignments in subroutine if target_node is None: - if any( - isinstance(nd.value, ast.Tuple) - for nd in ast.walk(call_ast) - if isinstance(nd, ast.Return) - ): + return_nodes = [nd for nd in ast.walk(call_ast) if isinstance(nd, ast.Return)] + if any(isinstance(nd.value, ast.Tuple) for nd in return_nodes): raise GTScriptSyntaxError( "Only functions with a single return value can be used in expressions, including as call arguments. " "Please assign the function results to symbols first." ) - target_node = ast.Name( - ctx=ast.Store(), - lineno=node.lineno, - col_offset=node.col_offset, - id=template_fmt.format(name="RETURN_VALUE"), - ) - assert isinstance(target_node, (ast.Name, ast.Tuple, ast.Subscript)) and isinstance( - target_node.ctx, ast.Store - ) + if len(return_nodes) > 0: + target_node = ast.Name( + ctx=ast.Store(), + lineno=node.lineno, + col_offset=node.col_offset, + id=template_fmt.format(name="RETURN_VALUE"), + ) + assert target_node is None or ( + isinstance(target_node, (ast.Name, ast.Tuple, ast.Subscript)) + and isinstance(target_node.ctx, ast.Store) + ) ReturnReplacer.apply(call_ast, target_node) # Add subroutine sources prepending the required arg assignments @@ -552,7 +555,7 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex col_offset=target_node.col_offset, elts=target_node.elts, ) - else: + elif isinstance(target_node, ast.Subscript): result_node = ast.Subscript( ctx=ast.Load(), lineno=target_node.lineno, @@ -560,6 +563,8 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex value=target_node.value, slice=target_node.slice, ) + else: # target_node is None + result_node = call_ast.body[0] # Add the temp_annotations and temp_init_values to the parent current_info = self.context[self.current_name]._gtscript_ @@ -574,8 +579,15 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex return result_node def visit_Expr(self, node: ast.Expr): - """Ignore pure string statements in callee.""" - if not isinstance(node.value, (ast.Constant, ast.Str)): + if ( + isinstance(node.value, ast.Call) + and gt_meta.get_qualified_name_from_node(node.value.func) not in gtscript.MATH_BUILTINS + ): + # Inline a function with no return value and then remove the current node + self.visit(node.value, target_node=None) + return None + elif not isinstance(node.value, (ast.Constant, ast.Str)): + # Ignore ure string statements in callee return super().visit(node.value) From 45363a318e21a593f188f789c3de6c23c091d346 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 25 Oct 2024 22:36:20 +0200 Subject: [PATCH 11/48] Allow assigning call arguments inside functions. --- .../cartesian/frontend/gtscript_frontend.py | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index bed8fb3418..ef7a6ae6b7 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -419,6 +419,12 @@ def visit_Assign(self, node: ast.Assign): else: return self.generic_visit(node) + def _get_sliced_symbol(self, node): + if isinstance(node, ast.Name): + return node.id + elif isinstance(node, ast.Subscript): + return self._get_sliced_symbol(node.value) + def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complexity too high call_name = gt_meta.get_qualified_name_from_node(node.func) @@ -476,10 +482,15 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex assigned_symbols = set() for target in assign_targets: - if not isinstance(target, ast.Name): - raise GTScriptSyntaxError(message="Unsupported assignment target.", loc=target) + if isinstance(target, ast.Subscript): + sliced_symbol = self._get_sliced_symbol(target) + if sliced_symbol not in call_args: + raise GTScriptSyntaxError(message="Unsupported assignment target.", loc=target) + else: + if not isinstance(target, ast.Name): + raise GTScriptSyntaxError(message="Unsupported assignment target.", loc=target) - assigned_symbols.add(target.id) + assigned_symbols.add(target.id) name_mapping = { name: value.id From 7015268c42525c7c1768e4645c8528bfc36b6df6 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 21 Nov 2024 15:46:54 +0100 Subject: [PATCH 12/48] Add inlining of constant function arguments. --- src/gt4py/cartesian/frontend/gtscript_frontend.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index ef7a6ae6b7..42710dec7d 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -472,6 +472,14 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex message="Invalid call signature", loc=nodes.Location.from_ast_node(node) ) from ex + # Inline constant function arguments + local_context = { + name: arg_node.value + for name, arg_node in call_args.items() + if isinstance(arg_node, ast.Constant) + } + ValueInliner.apply(call_ast, local_context) + # Rename local names in subroutine to avoid conflicts with caller context names try: assign_targets = gt_meta.collect_assign_targets(call_ast, allow_multiple_targets=False) From d151a7880518c340bb862e49aab962c701fe7720 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 22 Nov 2024 17:43:09 +0100 Subject: [PATCH 13/48] Inline constant function arguments recursively. --- .../cartesian/frontend/gtscript_frontend.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 42710dec7d..cb56c95727 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -441,15 +441,8 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex elif call_name not in self.context or not hasattr(self.context[call_name], "_gtscript_"): raise GTScriptSyntaxError("Unknown call", loc=nodes.Location.from_ast_node(node)) - # Recursively inline any possible nested subroutine call - call_info = self.context[call_name]._gtscript_ - call_ast = copy.deepcopy(call_info["ast"]) - self.current_name = call_name - CallInliner.apply( - call_ast, call_info["local_context"], call_stack={*self.call_stack, call_name} - ) - # Extract call arguments + call_info = self.context[call_name]._gtscript_ call_signature = call_info["api_signature"] arg_infos = {arg.name: arg.default for arg in call_signature} try: @@ -473,6 +466,7 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex ) from ex # Inline constant function arguments + call_ast = copy.deepcopy(call_info["ast"]) local_context = { name: arg_node.value for name, arg_node in call_args.items() @@ -480,6 +474,12 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex } ValueInliner.apply(call_ast, local_context) + # Recursively inline any possible nested subroutine call + self.current_name = call_name + CallInliner.apply( + call_ast, call_info["local_context"], call_stack={*self.call_stack, call_name} + ) + # Rename local names in subroutine to avoid conflicts with caller context names try: assign_targets = gt_meta.collect_assign_targets(call_ast, allow_multiple_targets=False) From 10c1bfbcec4fa75f14be5690431418dbffba8bfd Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 22 Nov 2024 23:13:45 +0100 Subject: [PATCH 14/48] Properly support for-loops around functions. --- src/gt4py/cartesian/frontend/gtscript_frontend.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index cb56c95727..b55e068d8e 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -2152,10 +2152,15 @@ def run(self, backend_name: str): ValueInliner.apply(main_func_node, context=local_context) + # unroll loops over data dimensions + # note(stubbiali): address the case of a function called within a for-loop + DataDimLoopUnroller.apply(main_func_node, context=local_context) + # Inline function calls CallInliner.apply(main_func_node, context=local_context) # unroll loops over data dimensions + # note(stubbiali): address the case of a for-loop inside a function DataDimLoopUnroller.apply(main_func_node, context=local_context) # Evaluate and inline compile-time conditionals From c811ba54089b4fdbcd5f873ac0a6d020c17bd9b1 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 19 Dec 2024 21:33:03 +0100 Subject: [PATCH 15/48] Improve support for global tables in numpy generated code. --- src/gt4py/cartesian/gtc/numpy/npir_codegen.py | 22 +++++++++++++------ 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py index 40bd75e534..c1dc447a1b 100644 --- a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py +++ b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py @@ -79,22 +79,29 @@ def _make_slice_access( """\ class Field: def __init__(self, field, offsets: Tuple[int, ...], dimensions: Tuple[bool, bool, bool]): - ii = iter(range(3)) - self.idx_to_data = tuple( - [next(ii) if has_dim else None for has_dim in dimensions] - + list(range(sum(dimensions), len(field.shape))) - ) + self.is_global_table = all(not has_dim for has_dim in dimensions) + + if self.is_global_table: + self.idx_to_data = tuple(i for i in range(field.ndim)) + self.offsets = (0,) * field.ndim + else: + self.idx_to_data = tuple( + [i if has_dim else None for i, has_dim in enumerate(dimensions)] + + list(range(sum(dimensions), field.ndim)) + ) + self.offsets = offsets shape = [field.shape[i] if i is not None else 1 for i in self.idx_to_data] self.field_view = np.reshape(field.data, shape).view(np.ndarray) - self.offsets = offsets - @classmethod def empty(cls, shape, dtype, offset): return cls(np.empty(shape, dtype=dtype), offset, (True, True, True)) def shim_key(self, key): + if self.is_global_table: + return key + new_args = [] if not isinstance(key, tuple): key = (key, ) @@ -134,6 +141,7 @@ def __getitem__(self, key): return self.field_view.__getitem__(self.shim_key(key)) def __setitem__(self, key, value): + assert not self.is_global_table return self.field_view.__setitem__(self.shim_key(key), value) """ ) From 878fd3586e18702f8ca21f714e032cb8cbb9d6db Mon Sep 17 00:00:00 2001 From: stubbiali Date: Mon, 27 Jan 2025 10:37:36 +0100 Subject: [PATCH 16/48] Fully inline constant arguments. --- src/gt4py/cartesian/frontend/gtscript_frontend.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 182e01b0a6..43b5a49f26 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -466,7 +466,7 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex loc=nodes.Location.from_ast_node(node), ) from ex - # Inline constant function arguments + # Inline constant arguments call_ast = copy.deepcopy(call_info["ast"]) local_context = { name: arg_node.value @@ -541,7 +541,8 @@ def visit_Call(self, node: ast.Call, *, target_node=None): # Cyclomatic complex # Add subroutine sources prepending the required arg assignments inlined_stmts = [] for arg_name, arg_value in call_args.items(): - if arg_name not in name_mapping: + # note(stubbiali): filter out constant arguments (which have been previously inlined) + if arg_name not in name_mapping and not isinstance(arg_value, ast.Constant): inlined_stmts.append( ast.Assign( lineno=node.lineno, From 4f25fa30e83e6dd4d94967a44eeb11f51b08dd00 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Tue, 4 Feb 2025 23:14:19 +0100 Subject: [PATCH 17/48] Prototype implementation of reductions. --- .../cartesian/frontend/gtscript_frontend.py | 123 ++++++++++++++---- src/gt4py/cartesian/gtscript.py | 17 ++- 2 files changed, 113 insertions(+), 27 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index db594f4053..34fc07bcc5 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -407,10 +407,9 @@ def visit_While(self, node: ast.While): return node def visit_Assign(self, node: ast.Assign): - if ( - isinstance(node.value, ast.Call) - and gt_meta.get_qualified_name_from_node(node.value.func) not in gtscript.MATH_BUILTINS - ): + if isinstance(node.value, ast.Call) and gt_meta.get_qualified_name_from_node( + node.value.func + ) not in gtscript.MATH_BUILTINS.union(gtscript.REDUCTION_BUILTINS): assert len(node.targets) == 1 self.visit(node.value, target_node=node.targets[0]) # This node can be now removed since the trivial assignment has been already done @@ -646,7 +645,7 @@ def visit_If(self, node: ast.If): return node if node else None -class DataDimLoopIndexReplacer(ast.NodeTransformer): +class LoopIndexReplacer(ast.NodeTransformer): def __init__(self, name: str, value: int) -> None: self.name = name self.value = value @@ -658,6 +657,94 @@ def visit_Name(self, node: ast.Name) -> Union[ast.Constant, ast.Name]: return node +def _get_loop_index_values( + node: Union[ast.Call, ast.List, ast.Tuple], context: dict +) -> Optional[Union[list, range]]: + if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "range": + range_args = [eval(ast.unparse(arg), context) for arg in node.args] + assert 1 <= len(range_args) <= 3 + if len(range_args) == 1: + start, stop, step = 0, *range_args, 1 + elif len(range_args) == 2: + start, stop, step = *range_args, 1 + else: + start, stop, step = range_args + index_values = range(start, stop, step) + elif isinstance(node, (ast.List, ast.Tuple)): + index_values = [eval(ast.unparse(elt), context) for elt in node.elts] + else: + index_values = None + return index_values + + +class ReductionUnroller(ast.NodeTransformer): + REDUCTION_OP_TO_AST_OP = {"add": ast.Add} + + @classmethod + def apply(cls, func_node: ast.FunctionDef, context: dict): + unroller = cls(context) + unroller(func_node) + + def __init__(self, context: dict) -> None: + self.context = context + + def __call__(self, func_node: ast.FunctionDef) -> None: + self.visit(func_node) + + def visit_Call(self, node: ast.Call) -> Union[ast.Call, ast.BinOp]: + if isinstance(node.func, ast.Name) and node.func.id == "reduce": + return self._unroll_reduction(node) + else: + return node + + def _unroll_reduction(self, node: ast.Call) -> ast.BinOp: + param_names = ["op", "generator", "initial"][len(args := node.args) :] + for kwarg in node.keywords: + if kwarg.arg in param_names: + args.append(kwarg.value) + else: + raise GTScriptSyntaxError(f"Reduce: unknown argument `{kwarg.arg}`.") + + if not 2 <= len(args) <= 3: + raise GTScriptSyntaxError("Reduce: the function takes 2 to 3 arguments.") + + if isinstance(args[0], ast.Name) and (op_id := args[0].id) in self.REDUCTION_OP_TO_AST_OP: + op = self.REDUCTION_OP_TO_AST_OP[op_id]() + else: + raise GTScriptSyntaxError("Reduce: invalid reduction operator.") + + if isinstance((generator_expr := args[1]), ast.GeneratorExp): + template_item = generator_expr.elt + index_name = generator_expr.generators[0].target.id + index_values = list( + _get_loop_index_values(generator_expr.generators[0].iter, self.context) + ) + else: + raise GTScriptSyntaxError("Reduce: second argument should be a generator expression.") + + initial_value = args[2] if len(node.args) == 3 else None + + return self._get_binary_node( + op, template_item, index_name, index_values, left=initial_value + ) + + def _get_binary_node(self, op, template_item, index_name, index_values, left=None) -> ast.BinOp: + if left is None: + assert len(index_values) > 1 + left = LoopIndexReplacer(index_name, index_values[0]).visit( + copy.deepcopy(template_item) + ) + index_values = index_values[1:] + + assert len(index_values) > 0 + if len(index_values) == 1: + right = LoopIndexReplacer(index_name, index_values[0]).visit(template_item) + else: + right = self._get_binary_node(op, template_item, index_name, index_values) + + return ast.BinOp(left=left, op=op, right=right) + + class DataDimLoopUnroller(ast.NodeTransformer): @classmethod def apply(cls, func_node: ast.FunctionDef, context: dict): @@ -668,38 +755,20 @@ def __init__(self, context): self.context = context self.prefix = "" - def __call__(self, func_node: ast.FunctionDef): + def __call__(self, func_node: ast.FunctionDef) -> None: self.visit(func_node) def visit_For(self, node: ast.For) -> Union[ast.For, list[ast.AST]]: super().generic_visit(node) - index_values: Optional[Union[list, range]] = None - if ( - isinstance(node.iter, ast.Call) - and isinstance(node.iter.func, ast.Name) - and node.iter.func.id == "range" - ): - range_args = [eval(ast.unparse(arg), self.context) for arg in node.iter.args] - assert 1 <= len(range_args) <= 3 - if len(range_args) == 1: - start, stop, step = 0, *range_args, 1 - elif len(range_args) == 2: - start, stop, step = *range_args, 1 - else: - start, stop, step = range_args - index_values = range(start, stop, step) - elif isinstance(node.iter, (ast.List, ast.Tuple)): - index_values = [eval(ast.unparse(elt), self.context) for elt in node.iter.elts] - - if index_values is not None: + if (index_values := _get_loop_index_values(node.iter, self.context)) is not None: assert isinstance(node.target, ast.Name) index_name = node.target.id new_body = [] for index_value in index_values: body = copy.deepcopy(node.body) - transformer = DataDimLoopIndexReplacer(index_name, index_value) + transformer = LoopIndexReplacer(index_name, index_value) new_body_item = [transformer.visit(stmt) for stmt in body] new_body += new_body_item @@ -2146,6 +2215,8 @@ def run(self, backend_name: str): ValueInliner.apply(main_func_node, context=local_context) + ReductionUnroller.apply(main_func_node, context=local_context) + # unroll loops over data dimensions # note(stubbiali): address the case of a function called within a for-loop DataDimLoopUnroller.apply(main_func_node, context=local_context) diff --git a/src/gt4py/cartesian/gtscript.py b/src/gt4py/cartesian/gtscript.py index 2697a96bde..fe494042e8 100644 --- a/src/gt4py/cartesian/gtscript.py +++ b/src/gt4py/cartesian/gtscript.py @@ -65,6 +65,8 @@ "f64", } +REDUCTION_BUILTINS = {"reduce", "add"} + builtins = { "I", "J", @@ -89,9 +91,10 @@ "compile_assert", "range", *MATH_BUILTINS, + *REDUCTION_BUILTINS, } -IGNORE_WHEN_INLINING = {*MATH_BUILTINS, "compile_assert", "range"} +IGNORE_WHEN_INLINING = {*MATH_BUILTINS, *REDUCTION_BUILTINS, "compile_assert", "range"} __all__ = [*list(builtins), "function", "stencil", "lazy_stencil"] @@ -919,3 +922,15 @@ def ceil(x): def trunc(x): """Return the Real value x truncated to an Integral (usually an integer)""" pass + + +# GTScript builtins: reductions +def reduce(op, generator, initial=None): + """Apply the binary operator `op` cumulatively to all elements of `generator` with + initial value `initial` (optional).""" + pass + + +def add(x, y): + """Placeholder for sum reduction operator.""" + pass From a75a0defd047bc96a737a6b73cb4f57b70887708 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 5 Mar 2025 21:47:47 +0100 Subject: [PATCH 18/48] Fix numpy codegen. --- src/gt4py/cartesian/gtc/numpy/npir_codegen.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py index c1dc447a1b..5dc558e082 100644 --- a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py +++ b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py @@ -85,10 +85,16 @@ def __init__(self, field, offsets: Tuple[int, ...], dimensions: Tuple[bool, bool self.idx_to_data = tuple(i for i in range(field.ndim)) self.offsets = (0,) * field.ndim else: - self.idx_to_data = tuple( - [i if has_dim else None for i, has_dim in enumerate(dimensions)] - + list(range(sum(dimensions), field.ndim)) - ) + idx = 0 + idx_to_data = [] + for has_dim in dimensions: + if has_dim: + idx_to_data.append(idx) + idx += 1 + else: + idx_to_data.append(None) + idx_to_data += list(range(idx, field.ndim)) + self.idx_to_data = tuple(idx_to_data) self.offsets = offsets shape = [field.shape[i] if i is not None else 1 for i in self.idx_to_data] From 01e7555e0b6810357a2e85f38720ccedb80d4e0f Mon Sep 17 00:00:00 2001 From: stubbiali Date: Tue, 20 May 2025 08:49:43 +0200 Subject: [PATCH 19/48] Fix "TypeError: : cannot pickle 'PyCapsule' object". --- src/gt4py/cartesian/utils/meta.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/utils/meta.py b/src/gt4py/cartesian/utils/meta.py index 3f02ecce51..fbdc5acc7a 100644 --- a/src/gt4py/cartesian/utils/meta.py +++ b/src/gt4py/cartesian/utils/meta.py @@ -294,7 +294,7 @@ def apply(cls, ast_root, context, default=None): return result def __init__(self, context: dict): - self.context = copy.deepcopy(context) + self.context = {**context} def visit_Name(self, node): return self.context[node.id] From 12309205f6051744e16cb9ea7c4dab2adebdd361 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 21 May 2025 15:27:58 +0200 Subject: [PATCH 20/48] Add custom AST nodes. --- .../cartesian/frontend/gtscript_frontend.py | 233 +++++++----------- 1 file changed, 94 insertions(+), 139 deletions(-) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 34fc07bcc5..78b880645c 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -16,7 +16,20 @@ import time import types import warnings -from typing import Any, Dict, Final, List, Literal, Optional, Sequence, Set, Tuple, Type, Union +from typing import ( + Any, + ClassVar, + Dict, + Final, + List, + Literal, + Optional, + Sequence, + Set, + Tuple, + Type, + Union, +) import numpy as np @@ -25,6 +38,7 @@ from gt4py.cartesian.frontend.defir_to_gtir import DefIRToGTIR, UnrollVectorAssignments from gt4py.cartesian.gtc import utils as gtc_utils from gt4py.cartesian.utils import NOTHING, meta as gt_meta +from gt4py.eve import datamodels as gt_datamodels from .base import Frontend, register from .exceptions import ( @@ -341,6 +355,77 @@ def visit_Return(self, node: ast.Return, *, target_node: ast.AST) -> ast.Assign: ) +@gt_datamodels.datamodel(frozen=True) +class ForIndex(ast.AST): + name: str + + +@gt_datamodels.datamodel(frozen=True) +class ForIndexTransformer(ast.NodeTransformer): + name: str + + def visit_Name(self, node: ast.Name) -> Union[ForIndex, ast.Name]: + super().generic_visit(node) + return ForIndex(self.name) if node.id == self.name else node + + +@gt_datamodels.datamodel +class For(ast.AST): + index_name: str + index_values: Optional[range] + body: list[ast.AST] + _fields: ClassVar[tuple[str, ...]] = ("body",) + + def __post_init__(self) -> None: + transformer = ForIndexTransformer(self.index_name) + self.body = [transformer.visit(stmt) for stmt in self.body] + + +@gt_datamodels.datamodel(frozen=True) +class ForTransformer(ast.NodeTransformer): + context: dict + + @classmethod + def apply(cls, func_node: ast.FunctionDef, context: dict): + unroller = cls(context) + unroller(func_node) + + def __call__(self, func_node: ast.FunctionDef) -> None: + self.visit(func_node) + + def visit_For(self, node: Union[ast.For, For]) -> For: + super().generic_visit(node) + if isinstance(node, ast.For): + assert isinstance(node.target, ast.Name) + + if ( + isinstance(node.iter, ast.Call) + and isinstance(node.iter.func, ast.Name) + and node.iter.func.id == "range" + ): + range_args = [eval(ast.unparse(arg), self.context) for arg in node.iter.args] + assert 1 <= len(range_args) <= 3 + if len(range_args) == 1: + start, stop, step = 0, *range_args, 1 + elif len(range_args) == 2: + start, stop, step = *range_args, 1 + else: + start, stop, step = range_args + index_values = range(start, stop, step) + else: + raise GTScriptSyntaxError( + "For-loop index values can only be specified using range()." + ) + + return For( + index_name=node.target.id, + index_values=index_values, + body=[self.visit(item) for item in node.body], + ) + else: + return node + + class CallInliner(ast.NodeTransformer): """Inlines calls to gtscript.function calls. @@ -406,6 +491,10 @@ def visit_While(self, node: ast.While): node.body = self._process_stmts(node.body) return node + def visit_For(self, node: Union[ast.For, For]): + node.body = self._process_stmts(node.body) + return node + def visit_Assign(self, node: ast.Assign): if isinstance(node.value, ast.Call) and gt_meta.get_qualified_name_from_node( node.value.func @@ -645,138 +734,6 @@ def visit_If(self, node: ast.If): return node if node else None -class LoopIndexReplacer(ast.NodeTransformer): - def __init__(self, name: str, value: int) -> None: - self.name = name - self.value = value - - def visit_Name(self, node: ast.Name) -> Union[ast.Constant, ast.Name]: - if node.id == self.name: - return ast.Constant(self.value, ctx=node.ctx) - else: - return node - - -def _get_loop_index_values( - node: Union[ast.Call, ast.List, ast.Tuple], context: dict -) -> Optional[Union[list, range]]: - if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "range": - range_args = [eval(ast.unparse(arg), context) for arg in node.args] - assert 1 <= len(range_args) <= 3 - if len(range_args) == 1: - start, stop, step = 0, *range_args, 1 - elif len(range_args) == 2: - start, stop, step = *range_args, 1 - else: - start, stop, step = range_args - index_values = range(start, stop, step) - elif isinstance(node, (ast.List, ast.Tuple)): - index_values = [eval(ast.unparse(elt), context) for elt in node.elts] - else: - index_values = None - return index_values - - -class ReductionUnroller(ast.NodeTransformer): - REDUCTION_OP_TO_AST_OP = {"add": ast.Add} - - @classmethod - def apply(cls, func_node: ast.FunctionDef, context: dict): - unroller = cls(context) - unroller(func_node) - - def __init__(self, context: dict) -> None: - self.context = context - - def __call__(self, func_node: ast.FunctionDef) -> None: - self.visit(func_node) - - def visit_Call(self, node: ast.Call) -> Union[ast.Call, ast.BinOp]: - if isinstance(node.func, ast.Name) and node.func.id == "reduce": - return self._unroll_reduction(node) - else: - return node - - def _unroll_reduction(self, node: ast.Call) -> ast.BinOp: - param_names = ["op", "generator", "initial"][len(args := node.args) :] - for kwarg in node.keywords: - if kwarg.arg in param_names: - args.append(kwarg.value) - else: - raise GTScriptSyntaxError(f"Reduce: unknown argument `{kwarg.arg}`.") - - if not 2 <= len(args) <= 3: - raise GTScriptSyntaxError("Reduce: the function takes 2 to 3 arguments.") - - if isinstance(args[0], ast.Name) and (op_id := args[0].id) in self.REDUCTION_OP_TO_AST_OP: - op = self.REDUCTION_OP_TO_AST_OP[op_id]() - else: - raise GTScriptSyntaxError("Reduce: invalid reduction operator.") - - if isinstance((generator_expr := args[1]), ast.GeneratorExp): - template_item = generator_expr.elt - index_name = generator_expr.generators[0].target.id - index_values = list( - _get_loop_index_values(generator_expr.generators[0].iter, self.context) - ) - else: - raise GTScriptSyntaxError("Reduce: second argument should be a generator expression.") - - initial_value = args[2] if len(node.args) == 3 else None - - return self._get_binary_node( - op, template_item, index_name, index_values, left=initial_value - ) - - def _get_binary_node(self, op, template_item, index_name, index_values, left=None) -> ast.BinOp: - if left is None: - assert len(index_values) > 1 - left = LoopIndexReplacer(index_name, index_values[0]).visit( - copy.deepcopy(template_item) - ) - index_values = index_values[1:] - - assert len(index_values) > 0 - if len(index_values) == 1: - right = LoopIndexReplacer(index_name, index_values[0]).visit(template_item) - else: - right = self._get_binary_node(op, template_item, index_name, index_values) - - return ast.BinOp(left=left, op=op, right=right) - - -class DataDimLoopUnroller(ast.NodeTransformer): - @classmethod - def apply(cls, func_node: ast.FunctionDef, context: dict): - unroller = cls(context) - unroller(func_node) - - def __init__(self, context): - self.context = context - self.prefix = "" - - def __call__(self, func_node: ast.FunctionDef) -> None: - self.visit(func_node) - - def visit_For(self, node: ast.For) -> Union[ast.For, list[ast.AST]]: - super().generic_visit(node) - - if (index_values := _get_loop_index_values(node.iter, self.context)) is not None: - assert isinstance(node.target, ast.Name) - index_name = node.target.id - - new_body = [] - for index_value in index_values: - body = copy.deepcopy(node.body) - transformer = LoopIndexReplacer(index_name, index_value) - new_body_item = [transformer.visit(stmt) for stmt in body] - new_body += new_body_item - - return new_body - else: - return node - - def _make_temp_decls( descriptors: Dict[str, gtscript._FieldDescriptor], ) -> Dict[str, nodes.FieldDecl]: @@ -2215,18 +2172,16 @@ def run(self, backend_name: str): ValueInliner.apply(main_func_node, context=local_context) - ReductionUnroller.apply(main_func_node, context=local_context) - - # unroll loops over data dimensions + # Insert custom nodes for for-loops # note(stubbiali): address the case of a function called within a for-loop - DataDimLoopUnroller.apply(main_func_node, context=local_context) + ForTransformer.apply(main_func_node, context=local_context) # Inline function calls CallInliner.apply(main_func_node, context=local_context) - # unroll loops over data dimensions + # Insert custom nodes for for-loops # note(stubbiali): address the case of a for-loop inside a function - DataDimLoopUnroller.apply(main_func_node, context=local_context) + ForTransformer.apply(main_func_node, context=local_context) # Evaluate and inline compile-time conditionals CompiledIfInliner.apply(main_func_node, context=local_context, stencil_name=self.main_name) From a5398fbeb61d28a4c1b0ecda2182f7562bc715d2 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 21 May 2025 15:37:16 +0200 Subject: [PATCH 21/48] Add For and ForIndex defir nodes. --- .../cartesian/frontend/gtscript_frontend.py | 31 +++++++++++++++++++ src/gt4py/cartesian/frontend/nodes.py | 15 +++++++++ 2 files changed, 46 insertions(+) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 78b880645c..bb85272a72 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -1502,6 +1502,37 @@ def visit_While(self, node: ast.While) -> list: return result + def visit_ForIndex(self, node: ForIndex) -> nodes.ForIndex: + return nodes.ForIndex(name=node.name) + + def visit_For(self, node: For) -> list: + assert isinstance(node, For) + + loc = nodes.Location.from_ast_node(node) + + self.decls_stack.append([]) + stmts = gt_utils.flatten([self.visit(stmt) for stmt in node.body]) + assert all(isinstance(item, nodes.Statement) for item in stmts) + + result = [ + nodes.For( + index=nodes.ForIndex(name=node.index_name), + iter_start=node.index_values.start, + iter_stop=node.index_values.stop, + iter_step=node.index_values.step, + body=nodes.BlockStmt(stmts=stmts, loc=loc), + loc=nodes.Location.from_ast_node(node), + ) + ] + + if len(self.decls_stack) == 1: + result.extend(self.decls_stack.pop()) + elif len(self.decls_stack) > 1: + self.decls_stack[-2].extend(self.decls_stack[-1]) + self.decls_stack.pop() + + return result + def visit_Call(self, node: ast.Call): native_fcn = nodes.NativeFunction.PYTHON_SYMBOL_TO_IR_OP[node.func.id] diff --git a/src/gt4py/cartesian/frontend/nodes.py b/src/gt4py/cartesian/frontend/nodes.py index ab610e81e4..87e64c05d6 100644 --- a/src/gt4py/cartesian/frontend/nodes.py +++ b/src/gt4py/cartesian/frontend/nodes.py @@ -376,6 +376,11 @@ class AxisIndex(Expr): data_type = attribute(of=DataType, default=DataType.INT32) +@attribclass +class ForIndex(Expr): + name = attribute(of=str) + + @enum.unique class NativeFunction(enum.Enum): ABS = enum.auto() @@ -654,6 +659,16 @@ class While(Statement): loc = attribute(of=Location, optional=True) +@attribclass +class For(Statement): + index = attribute(of=ForIndex) + iter_start = attribute(of=int) + iter_stop = attribute(of=int) + iter_step = attribute(of=int) + body = attribute(of=BlockStmt) + loc = attribute(of=Location, optional=None) + + # ---- IR: computations ---- @enum.unique class IterationOrder(enum.Enum): From 1e65bd519f176058a877b904fbbdeda239c33ab7 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 21 May 2025 16:07:57 +0200 Subject: [PATCH 22/48] Add custom For and ForIndex gtir nodes. --- src/gt4py/cartesian/frontend/defir_to_gtir.py | 15 +++++++++++++++ src/gt4py/cartesian/gtc/common.py | 12 ++++++++++++ src/gt4py/cartesian/gtc/gtir.py | 9 +++++++++ 3 files changed, 36 insertions(+) diff --git a/src/gt4py/cartesian/frontend/defir_to_gtir.py b/src/gt4py/cartesian/frontend/defir_to_gtir.py index ede29efc6f..513ade8f9f 100644 --- a/src/gt4py/cartesian/frontend/defir_to_gtir.py +++ b/src/gt4py/cartesian/frontend/defir_to_gtir.py @@ -36,6 +36,8 @@ Expr, FieldDecl, FieldRef, + For, + ForIndex, HorizontalIf, If, IterationOrder, @@ -528,6 +530,19 @@ def visit_While(self, node: While) -> gtir.While: loc=location_to_source_location(node.loc), ) + def visit_ForIndex(self, node: ForIndex) -> gtir.ForIndex: + return gtir.ForIndex(name=node.name) + + def visit_For(self, node: For) -> gtir.For: + return gtir.For( + index=self.visit(node.index), + iter_start=node.iter_start, + iter_stop=node.iter_stop, + iter_step=node.iter_step, + body=self.visit(node.body), + loc=location_to_source_location(node.loc), + ) + def visit_VarRef(self, node: VarRef, **kwargs): return gtir.ScalarAccess(name=node.name, loc=location_to_source_location(node.loc)) diff --git a/src/gt4py/cartesian/gtc/common.py b/src/gt4py/cartesian/gtc/common.py index ec69bc1002..db34070389 100644 --- a/src/gt4py/cartesian/gtc/common.py +++ b/src/gt4py/cartesian/gtc/common.py @@ -399,6 +399,18 @@ def condition_is_boolean(self, attribute: datamodels.Attribute, value: Expr) -> verify_condition_is_boolean(self, value) +class ForIndex(eve.GenericNode): + name: str + + +class For(eve.GenericNode, Generic[StmtT]): + index: ForIndex + iter_start: int + iter_stop: int + iter_step: int + body: List[StmtT] + + class AssignStmt(eve.GenericNode, Generic[TargetT, ExprT]): left: TargetT right: ExprT diff --git a/src/gt4py/cartesian/gtc/gtir.py b/src/gt4py/cartesian/gtc/gtir.py index 0ee4f7ebe1..053fda5309 100644 --- a/src/gt4py/cartesian/gtc/gtir.py +++ b/src/gt4py/cartesian/gtc/gtir.py @@ -151,6 +151,15 @@ def _no_write_and_read_with_horizontal_offset_all( raise ValueError(f"Illegal write and read with horizontal offset detected for {names}.") +class ForIndex(common.ForIndex, Expr): + kind: common.ExprKind = common.ExprKind.SCALAR + dtype: common.DataType = common.DataType.INT64 + + +class For(common.For[Stmt], Stmt): + pass + + class UnaryOp(common.UnaryOp[Expr], Expr): pass From d13a8abb03ec6841b517a9ae5c668822cc390389 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 21 May 2025 22:16:50 +0200 Subject: [PATCH 23/48] Refactor common ForIndex. --- src/gt4py/cartesian/gtc/common.py | 4 +++- src/gt4py/cartesian/gtc/gtir.py | 3 +-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/src/gt4py/cartesian/gtc/common.py b/src/gt4py/cartesian/gtc/common.py index db34070389..a9060c4d38 100644 --- a/src/gt4py/cartesian/gtc/common.py +++ b/src/gt4py/cartesian/gtc/common.py @@ -399,8 +399,10 @@ def condition_is_boolean(self, attribute: datamodels.Attribute, value: Expr) -> verify_condition_is_boolean(self, value) -class ForIndex(eve.GenericNode): +class ForIndex(eve.GenericNode, Expr): name: str + kind: ExprKind = ExprKind.SCALAR + dtype: DataType = DataType.INT64 class For(eve.GenericNode, Generic[StmtT]): diff --git a/src/gt4py/cartesian/gtc/gtir.py b/src/gt4py/cartesian/gtc/gtir.py index 053fda5309..d7ec40ea96 100644 --- a/src/gt4py/cartesian/gtc/gtir.py +++ b/src/gt4py/cartesian/gtc/gtir.py @@ -152,8 +152,7 @@ def _no_write_and_read_with_horizontal_offset_all( class ForIndex(common.ForIndex, Expr): - kind: common.ExprKind = common.ExprKind.SCALAR - dtype: common.DataType = common.DataType.INT64 + pass class For(common.For[Stmt], Stmt): From c8dd83127891fd093308f341156c7e82c7a4b3d2 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 21 May 2025 22:17:15 +0200 Subject: [PATCH 24/48] Add For and ForIndex oir nodes. --- src/gt4py/cartesian/gtc/gtir_to_oir.py | 18 ++++++++++++++++++ src/gt4py/cartesian/gtc/oir.py | 8 ++++++++ 2 files changed, 26 insertions(+) diff --git a/src/gt4py/cartesian/gtc/gtir_to_oir.py b/src/gt4py/cartesian/gtc/gtir_to_oir.py index 96f8077ec4..4f56343c2a 100644 --- a/src/gt4py/cartesian/gtc/gtir_to_oir.py +++ b/src/gt4py/cartesian/gtc/gtir_to_oir.py @@ -139,6 +139,24 @@ def visit_While(self, node: gtir.While, **kwargs: Any) -> oir.While: condition: oir.Expr = self.visit(node.cond) return oir.While(cond=condition, body=body, loc=node.loc) + def visit_ForIndex(self, node: gtir.ForIndex, **kwargs: Any) -> oir.ForIndex: + return oir.ForIndex(name=node.name) + + def visit_For(self, node: gtir.For, **kwargs: Any) -> oir.For: + body: List[oir.Stmt] = [] + for statement in node.body: + oir_statement = self.visit(statement, **kwargs) + body.extend(utils.flatten_list(utils.listify(oir_statement))) + + return oir.For( + index=self.visit(node.index), + iter_start=node.iter_start, + iter_stop=node.iter_stop, + iter_step=node.iter_step, + body=body, + loc=node.loc, + ) + def visit_FieldIfStmt( self, node: gtir.FieldIfStmt, diff --git a/src/gt4py/cartesian/gtc/oir.py b/src/gt4py/cartesian/gtc/oir.py index 9f24db6e48..2e59b54b35 100644 --- a/src/gt4py/cartesian/gtc/oir.py +++ b/src/gt4py/cartesian/gtc/oir.py @@ -100,6 +100,14 @@ class While(common.While[Stmt, Expr], Stmt): pass +class ForIndex(common.ForIndex, Expr): + pass + + +class For(common.For[Stmt], Stmt): + pass + + class Decl(LocNode): name: eve.Coerced[eve.SymbolName] dtype: common.DataType From 6ed7ddd3efed2e41fdce6d83087b54972b2a94e5 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 21 May 2025 22:17:50 +0200 Subject: [PATCH 25/48] Add For and ForIndex gtcpp nodes. --- src/gt4py/cartesian/gtc/gtcpp/gtcpp.py | 8 ++++++++ src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py | 12 ++++++++++++ 2 files changed, 20 insertions(+) diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp.py index 5ca766c272..9d74e9eb57 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp.py @@ -72,6 +72,14 @@ class While(common.While[Stmt, Expr], Stmt): pass +class ForIndex(common.ForIndex, Expr): + pass + + +class For(common.For[Stmt], Stmt): + pass + + class UnaryOp(common.UnaryOp[Expr], Expr): pass diff --git a/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py b/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py index 0d5b1517c5..bcef02ad4f 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py +++ b/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py @@ -286,6 +286,18 @@ def visit_While(self, node: oir.While, **kwargs: Any) -> gtcpp.While: cond=self.visit(node.cond, **kwargs), body=self.visit(node.body, **kwargs) ) + def visit_ForIndex(self, node: oir.ForIndex, **kwargs: Any) -> gtcpp.ForIndex: + return gtcpp.ForIndex(name=node.name) + + def visit_For(self, node: common.For, **kwargs: Any) -> gtcpp.For: + return gtcpp.For( + index=self.visit(node.index, **kwargs), + iter_start=node.iter_start, + iter_stop=node.iter_stop, + iter_step=node.iter_step, + body=self.visit(node.body, **kwargs), + ) + def visit_HorizontalExecution( self, node: oir.HorizontalExecution, From 8570ce195ef2ac35a586d5448df0574b4238a5fe Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 21 May 2025 22:18:06 +0200 Subject: [PATCH 26/48] Support For and ForIndex in codegen. --- src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py index d9790249d9..459ef986ce 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py @@ -100,7 +100,7 @@ def visit_AccessorRef( temp = temp_decls[accessor_ref.name] data_index = "+".join( [ - f"{self.visit(index, in_data_index=True, **kwargs)}*{int(np.prod(temp.data_dims[i+1:], initial=1))}" + f"{self.visit(index, in_data_index=True, **kwargs)}*{int(np.prod(temp.data_dims[i + 1 :], initial=1))}" for i, index in enumerate(accessor_ref.data_index) ] ) @@ -252,6 +252,14 @@ def visit_Temporary(self, node: gtcpp.Temporary, **kwargs: Any) -> str: While = as_mako("while(${cond}) {${''.join(body)}}") + ForIndex = as_mako("${name}") + For = as_mako( + "for(std::size_t ${_this_node.index.name}=${iter_start}; " + "${_this_node.index.name}${'<' if _this_node.iter_step > 0 else '>'}${iter_stop}; " + "${_this_node.index.name}+=(${iter_step})) " + "{${''.join(body)}}" + ) + BlockStmt = as_mako("{${''.join(body)}}") def visit_GTComputationCall( From de465c65b8f56215f5613cc6858e68e74bcfcaa0 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 22 May 2025 09:46:22 +0200 Subject: [PATCH 27/48] Make For and ForIndex deep-copyable. --- src/gt4py/cartesian/frontend/gtscript_frontend.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index bb85272a72..20e77cde5a 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -359,6 +359,9 @@ def visit_Return(self, node: ast.Return, *, target_node: ast.AST) -> ast.Assign: class ForIndex(ast.AST): name: str + def __deepcopy__(self, memo: dict) -> "ForIndex": + return self + @gt_datamodels.datamodel(frozen=True) class ForIndexTransformer(ast.NodeTransformer): @@ -380,6 +383,13 @@ def __post_init__(self) -> None: transformer = ForIndexTransformer(self.index_name) self.body = [transformer.visit(stmt) for stmt in self.body] + def __deepcopy__(self, memo: dict) -> "For": + return For( + index_name=self.index_name, + index_values=self.index_values, + body=copy.deepcopy(self.body), + ) + @gt_datamodels.datamodel(frozen=True) class ForTransformer(ast.NodeTransformer): From cf0d34b778bf86695804345136375867cd4b3316 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 22 May 2025 10:36:55 +0200 Subject: [PATCH 28/48] index -> index_name --- src/gt4py/cartesian/frontend/defir_to_gtir.py | 2 +- .../cartesian/frontend/gtscript_frontend.py | 27 +++++++++++++------ src/gt4py/cartesian/frontend/nodes.py | 2 +- src/gt4py/cartesian/gtc/common.py | 2 +- .../cartesian/gtc/gtcpp/gtcpp_codegen.py | 6 ++--- src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py | 2 +- src/gt4py/cartesian/gtc/gtir_to_oir.py | 2 +- 7 files changed, 27 insertions(+), 16 deletions(-) diff --git a/src/gt4py/cartesian/frontend/defir_to_gtir.py b/src/gt4py/cartesian/frontend/defir_to_gtir.py index 513ade8f9f..c1683fd20a 100644 --- a/src/gt4py/cartesian/frontend/defir_to_gtir.py +++ b/src/gt4py/cartesian/frontend/defir_to_gtir.py @@ -535,7 +535,7 @@ def visit_ForIndex(self, node: ForIndex) -> gtir.ForIndex: def visit_For(self, node: For) -> gtir.For: return gtir.For( - index=self.visit(node.index), + index_name=node.index_name, iter_start=node.iter_start, iter_stop=node.iter_stop, iter_step=node.iter_step, diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 20e77cde5a..846e088109 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -375,8 +375,12 @@ def visit_Name(self, node: ast.Name) -> Union[ForIndex, ast.Name]: @gt_datamodels.datamodel class For(ast.AST): index_name: str - index_values: Optional[range] + iter_start: int + iter_stop: int + iter_step: int body: list[ast.AST] + lineno: Optional[int] = None + col_offset: Optional[int] = None _fields: ClassVar[tuple[str, ...]] = ("body",) def __post_init__(self) -> None: @@ -386,8 +390,12 @@ def __post_init__(self) -> None: def __deepcopy__(self, memo: dict) -> "For": return For( index_name=self.index_name, - index_values=self.index_values, + iter_start=self.iter_start, + iter_stop=self.iter_stop, + iter_step=self.iter_step, body=copy.deepcopy(self.body), + lineno=self.lineno, + col_offset=self.col_offset, ) @@ -421,7 +429,6 @@ def visit_For(self, node: Union[ast.For, For]) -> For: start, stop, step = *range_args, 1 else: start, stop, step = range_args - index_values = range(start, stop, step) else: raise GTScriptSyntaxError( "For-loop index values can only be specified using range()." @@ -429,8 +436,12 @@ def visit_For(self, node: Union[ast.For, For]) -> For: return For( index_name=node.target.id, - index_values=index_values, + iter_start=start, + iter_stop=stop, + iter_step=step, body=[self.visit(item) for item in node.body], + lineno=node.lineno, + col_offset=node.col_offset, ) else: return node @@ -1526,10 +1537,10 @@ def visit_For(self, node: For) -> list: result = [ nodes.For( - index=nodes.ForIndex(name=node.index_name), - iter_start=node.index_values.start, - iter_stop=node.index_values.stop, - iter_step=node.index_values.step, + index_name=node.index_name, + iter_start=node.iter_start, + iter_stop=node.iter_stop, + iter_step=node.iter_step, body=nodes.BlockStmt(stmts=stmts, loc=loc), loc=nodes.Location.from_ast_node(node), ) diff --git a/src/gt4py/cartesian/frontend/nodes.py b/src/gt4py/cartesian/frontend/nodes.py index 87e64c05d6..3b4473d715 100644 --- a/src/gt4py/cartesian/frontend/nodes.py +++ b/src/gt4py/cartesian/frontend/nodes.py @@ -661,7 +661,7 @@ class While(Statement): @attribclass class For(Statement): - index = attribute(of=ForIndex) + index_name = attribute(of=str) iter_start = attribute(of=int) iter_stop = attribute(of=int) iter_step = attribute(of=int) diff --git a/src/gt4py/cartesian/gtc/common.py b/src/gt4py/cartesian/gtc/common.py index a9060c4d38..4ccc57bfd0 100644 --- a/src/gt4py/cartesian/gtc/common.py +++ b/src/gt4py/cartesian/gtc/common.py @@ -406,7 +406,7 @@ class ForIndex(eve.GenericNode, Expr): class For(eve.GenericNode, Generic[StmtT]): - index: ForIndex + index_name: str iter_start: int iter_stop: int iter_step: int diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py index 459ef986ce..7b9c525d46 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py @@ -254,9 +254,9 @@ def visit_Temporary(self, node: gtcpp.Temporary, **kwargs: Any) -> str: ForIndex = as_mako("${name}") For = as_mako( - "for(std::size_t ${_this_node.index.name}=${iter_start}; " - "${_this_node.index.name}${'<' if _this_node.iter_step > 0 else '>'}${iter_stop}; " - "${_this_node.index.name}+=(${iter_step})) " + "for(std::size_t ${index_name}=${iter_start}; " + "${index_name}${'<' if _this_node.iter_step > 0 else '>'}${iter_stop}; " + "${index_name}+=(${iter_step})) " "{${''.join(body)}}" ) diff --git a/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py b/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py index bcef02ad4f..3094101f90 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py +++ b/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py @@ -291,7 +291,7 @@ def visit_ForIndex(self, node: oir.ForIndex, **kwargs: Any) -> gtcpp.ForIndex: def visit_For(self, node: common.For, **kwargs: Any) -> gtcpp.For: return gtcpp.For( - index=self.visit(node.index, **kwargs), + index_name=node.index_name, iter_start=node.iter_start, iter_stop=node.iter_stop, iter_step=node.iter_step, diff --git a/src/gt4py/cartesian/gtc/gtir_to_oir.py b/src/gt4py/cartesian/gtc/gtir_to_oir.py index 4f56343c2a..ce53878b91 100644 --- a/src/gt4py/cartesian/gtc/gtir_to_oir.py +++ b/src/gt4py/cartesian/gtc/gtir_to_oir.py @@ -149,7 +149,7 @@ def visit_For(self, node: gtir.For, **kwargs: Any) -> oir.For: body.extend(utils.flatten_list(utils.listify(oir_statement))) return oir.For( - index=self.visit(node.index), + index_name=node.index_name, iter_start=node.iter_start, iter_stop=node.iter_stop, iter_step=node.iter_step, From 9af5a5c1a9ca4de8bb638b753879a966d846a84b Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 23 May 2025 13:49:03 +0200 Subject: [PATCH 29/48] Add support for for-loops in dace. --- src/gt4py/cartesian/gtc/dace/daceir.py | 8 ++++++++ .../cartesian/gtc/dace/expansion/daceir_builder.py | 12 ++++++++++++ .../cartesian/gtc/dace/expansion/tasklet_codegen.py | 13 +++++++++++++ src/gt4py/cartesian/gtc/dace/utils.py | 3 +++ 4 files changed, 36 insertions(+) diff --git a/src/gt4py/cartesian/gtc/dace/daceir.py b/src/gt4py/cartesian/gtc/dace/daceir.py index 492a9598c5..fe9b8bd0d4 100644 --- a/src/gt4py/cartesian/gtc/dace/daceir.py +++ b/src/gt4py/cartesian/gtc/dace/daceir.py @@ -783,6 +783,14 @@ class While(common.While[Stmt, Expr], Stmt): pass +class ForIndex(common.ForIndex, Expr): + pass + + +class For(common.For[Stmt], Stmt): + pass + + class ScalarDecl(Decl): pass diff --git a/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py b/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py index e93a15debe..2450aa3fc1 100644 --- a/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py +++ b/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py @@ -394,6 +394,18 @@ def visit_While(self, node: oir.While, **kwargs: Any) -> dcir.While: body=self.visit(node.body, **kwargs), ) + def visit_ForIndex(self, node: oir.ForIndex, **kwargs: Any) -> dcir.ForIndex: + return dcir.ForIndex(name=node.name) + + def visit_For(self, node: oir.For, **kwargs: Any) -> dcir.For: + return dcir.For( + index_name=node.index_name, + iter_start=node.iter_start, + iter_stop=node.iter_stop, + iter_step=node.iter_step, + body=self.visit(node.body, **kwargs), + ) + def visit_Cast(self, node: oir.Cast, **kwargs: Any) -> dcir.Cast: return dcir.Cast(dtype=node.dtype, expr=self.visit(node.expr, **kwargs)) diff --git a/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py b/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py index 50aa695d39..2be188397c 100644 --- a/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py +++ b/src/gt4py/cartesian/gtc/dace/expansion/tasklet_codegen.py @@ -239,6 +239,19 @@ def visit_HorizontalRestriction(self, node: dcir.HorizontalRestriction, **kwargs def visit_While(self, node: dcir.While, **kwargs: Any) -> Any: return self._visit_conditional(cond=node.cond, body=node.body, keyword="while", **kwargs) + ForIndex = as_fmt("{name}") + + def visit_For(self, node: dcir.For, **kwargs: Any) -> str: + code = [ + f"for {node.index_name} in {range(node.iter_start, node.iter_stop, node.iter_step)}:", + *( + " " + line + for block in self.visit(node.body, **kwargs) + for line in block.split("\n") + ), + ] + return "\n".join(code) + def visit_HorizontalMask(self, node: common.HorizontalMask, **kwargs: Any) -> str: clauses: List[str] = [] diff --git a/src/gt4py/cartesian/gtc/dace/utils.py b/src/gt4py/cartesian/gtc/dace/utils.py index bd65861a49..e7ce97d609 100644 --- a/src/gt4py/cartesian/gtc/dace/utils.py +++ b/src/gt4py/cartesian/gtc/dace/utils.py @@ -189,6 +189,9 @@ def visit_MaskStmt(self, node: oir.MaskStmt, *, is_conditional=False, **kwargs): def visit_While(self, node: oir.While, *, is_conditional=False, **kwargs): self.generic_visit(node, is_conditional=True, **kwargs) + def visit_For(self, node: oir.For, *, is_conditional=False, **kwargs): + self.visit(node.body, is_conditional=False, **kwargs) + @staticmethod def _global_grid_subset( region: common.HorizontalMask, he_grid: dcir.GridSubset, offset: List[Optional[int]] From 011e3d46e741969253896efbe9c8adf5d295adf3 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Fri, 23 May 2025 14:15:35 +0200 Subject: [PATCH 30/48] Fix single precision. --- src/gt4py/cartesian/frontend/defir_to_gtir.py | 2 +- src/gt4py/cartesian/frontend/gtscript_frontend.py | 7 ++++++- src/gt4py/cartesian/frontend/nodes.py | 1 + src/gt4py/cartesian/gtc/common.py | 5 ++--- src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py | 2 +- src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py | 2 +- src/gt4py/cartesian/gtc/gtir_to_oir.py | 2 +- 7 files changed, 13 insertions(+), 8 deletions(-) diff --git a/src/gt4py/cartesian/frontend/defir_to_gtir.py b/src/gt4py/cartesian/frontend/defir_to_gtir.py index c1683fd20a..baa9aca649 100644 --- a/src/gt4py/cartesian/frontend/defir_to_gtir.py +++ b/src/gt4py/cartesian/frontend/defir_to_gtir.py @@ -531,7 +531,7 @@ def visit_While(self, node: While) -> gtir.While: ) def visit_ForIndex(self, node: ForIndex) -> gtir.ForIndex: - return gtir.ForIndex(name=node.name) + return gtir.ForIndex(name=node.name, dtype=common.DataType(node.data_type.value)) def visit_For(self, node: For) -> gtir.For: return gtir.For( diff --git a/src/gt4py/cartesian/frontend/gtscript_frontend.py b/src/gt4py/cartesian/frontend/gtscript_frontend.py index 846e088109..a827ff0f33 100644 --- a/src/gt4py/cartesian/frontend/gtscript_frontend.py +++ b/src/gt4py/cartesian/frontend/gtscript_frontend.py @@ -1524,7 +1524,12 @@ def visit_While(self, node: ast.While) -> list: return result def visit_ForIndex(self, node: ForIndex) -> nodes.ForIndex: - return nodes.ForIndex(name=node.name) + return nodes.ForIndex( + name=node.name, + data_type=nodes.DataType.from_dtype( + self.dtypes[int] if self.dtypes and int in self.dtypes else int + ), + ) def visit_For(self, node: For) -> list: assert isinstance(node, For) diff --git a/src/gt4py/cartesian/frontend/nodes.py b/src/gt4py/cartesian/frontend/nodes.py index 3b4473d715..a197a545e7 100644 --- a/src/gt4py/cartesian/frontend/nodes.py +++ b/src/gt4py/cartesian/frontend/nodes.py @@ -379,6 +379,7 @@ class AxisIndex(Expr): @attribclass class ForIndex(Expr): name = attribute(of=str) + data_type = attribute(of=DataType) @enum.unique diff --git a/src/gt4py/cartesian/gtc/common.py b/src/gt4py/cartesian/gtc/common.py index 4ccc57bfd0..1de88045a5 100644 --- a/src/gt4py/cartesian/gtc/common.py +++ b/src/gt4py/cartesian/gtc/common.py @@ -399,10 +399,9 @@ def condition_is_boolean(self, attribute: datamodels.Attribute, value: Expr) -> verify_condition_is_boolean(self, value) -class ForIndex(eve.GenericNode, Expr): +class ForIndex(eve.Node): name: str - kind: ExprKind = ExprKind.SCALAR - dtype: DataType = DataType.INT64 + dtype: DataType class For(eve.GenericNode, Generic[StmtT]): diff --git a/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py b/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py index 2450aa3fc1..b0c47d268f 100644 --- a/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py +++ b/src/gt4py/cartesian/gtc/dace/expansion/daceir_builder.py @@ -395,7 +395,7 @@ def visit_While(self, node: oir.While, **kwargs: Any) -> dcir.While: ) def visit_ForIndex(self, node: oir.ForIndex, **kwargs: Any) -> dcir.ForIndex: - return dcir.ForIndex(name=node.name) + return dcir.ForIndex(name=node.name, dtype=node.dtype) def visit_For(self, node: oir.For, **kwargs: Any) -> dcir.For: return dcir.For( diff --git a/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py b/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py index 3094101f90..945f8ea772 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py +++ b/src/gt4py/cartesian/gtc/gtcpp/oir_to_gtcpp.py @@ -287,7 +287,7 @@ def visit_While(self, node: oir.While, **kwargs: Any) -> gtcpp.While: ) def visit_ForIndex(self, node: oir.ForIndex, **kwargs: Any) -> gtcpp.ForIndex: - return gtcpp.ForIndex(name=node.name) + return gtcpp.ForIndex(name=node.name, dtype=node.dtype) def visit_For(self, node: common.For, **kwargs: Any) -> gtcpp.For: return gtcpp.For( diff --git a/src/gt4py/cartesian/gtc/gtir_to_oir.py b/src/gt4py/cartesian/gtc/gtir_to_oir.py index ce53878b91..fe81ebc7a1 100644 --- a/src/gt4py/cartesian/gtc/gtir_to_oir.py +++ b/src/gt4py/cartesian/gtc/gtir_to_oir.py @@ -140,7 +140,7 @@ def visit_While(self, node: gtir.While, **kwargs: Any) -> oir.While: return oir.While(cond=condition, body=body, loc=node.loc) def visit_ForIndex(self, node: gtir.ForIndex, **kwargs: Any) -> oir.ForIndex: - return oir.ForIndex(name=node.name) + return oir.ForIndex(name=node.name, dtype=node.dtype) def visit_For(self, node: gtir.For, **kwargs: Any) -> oir.For: body: List[oir.Stmt] = [] From 57d7f5f205570eac74be55894e7b44ea836ebfcf Mon Sep 17 00:00:00 2001 From: stubbiali Date: Tue, 7 Apr 2026 12:03:44 +0200 Subject: [PATCH 31/48] Fix cuda extra compile args --- src/gt4py/cartesian/config.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/config.py b/src/gt4py/cartesian/config.py index 5aa32506b7..2a4e55cc08 100644 --- a/src/gt4py/cartesian/config.py +++ b/src/gt4py/cartesian/config.py @@ -61,7 +61,14 @@ "gt_include_path": os.environ.get("GT_INCLUDE_PATH", GT_INCLUDE_PATH), "openmp_cppflags": os.environ.get("OPENMP_CPPFLAGS", "-fopenmp").split(), "openmp_ldflags": os.environ.get("OPENMP_LDFLAGS", "-fopenmp").split(), - "extra_compile_args": {"cxx": extra_compile_args, "cuda": extra_compile_args}, + "extra_compile_args": { + "cxx": extra_compile_args, + "cuda": [ + arg + for extra_compile_arg in extra_compile_args + for arg in f"--compiler-options {extra_compile_arg}".split(" ") + ], + }, "extra_link_args": extra_link_args, "parallel_jobs": multiprocessing.cpu_count(), "cpp_template_depth": os.environ.get("GT_CPP_TEMPLATE_DEPTH", GT_CPP_TEMPLATE_DEPTH), From 36c24fa260bd04a7385e79b2ac8b63329a6e94df Mon Sep 17 00:00:00 2001 From: stubbiali Date: Mon, 13 Apr 2026 11:18:28 +0200 Subject: [PATCH 32/48] Add support for for-loops in numpy backend --- src/gt4py/cartesian/gtc/numpy/npir.py | 8 +++++ src/gt4py/cartesian/gtc/numpy/npir_codegen.py | 24 +++++++++++++++ src/gt4py/cartesian/gtc/numpy/oir_to_npir.py | 29 ++++++++++++++++--- 3 files changed, 57 insertions(+), 4 deletions(-) diff --git a/src/gt4py/cartesian/gtc/numpy/npir.py b/src/gt4py/cartesian/gtc/numpy/npir.py index 36a17b7301..be11a0dc76 100644 --- a/src/gt4py/cartesian/gtc/numpy/npir.py +++ b/src/gt4py/cartesian/gtc/numpy/npir.py @@ -118,6 +118,10 @@ class VarKOffset(common.VariableKOffset[Expr]): pass +class ForIndex(common.ForIndex, Expr): + pass + + class FieldSlice(VectorLValue): name: eve.Coerced[eve.SymbolRef] i_offset: int @@ -181,6 +185,10 @@ class While(common.While[Stmt, Expr], Stmt): pass +class For(common.For[Stmt], Stmt): + pass + + # --- Control Flow --- class HorizontalBlock(common.LocNode, eve.SymbolTableTrait): body: List[Stmt] diff --git a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py index 5dc558e082..446e01ee96 100644 --- a/src/gt4py/cartesian/gtc/numpy/npir_codegen.py +++ b/src/gt4py/cartesian/gtc/numpy/npir_codegen.py @@ -179,6 +179,8 @@ def visit_TemporaryDecl( VarKOffset = as_fmt("lk + {k}") + ForIndex = as_fmt("{name}") + def visit_FieldSlice(self, node: npir.FieldSlice, **kwargs: Any) -> Union[str, Collection[str]]: k_offset = ( self.visit(node.k_offset, **kwargs) @@ -329,6 +331,28 @@ def visit_While(self, node: npir.While, **kwargs: Any) -> str: body.extend(stmt.split("\n")) return self.While.render(cond=cond, body=body) + For = as_jinja( + textwrap.dedent( + """\ + for {{ index_name }} in range({{ iter_start }}, {{ iter_stop }}, {{ iter_step }}): + {% for stmt in body %}{{ stmt }} + {% endfor %} + """ + ) + ) + + def visit_For(self, node: npir.For, **kwargs: Any) -> str: + body = [] + for stmt in self.visit(node.body, **kwargs): + body.extend(stmt.split("\n")) + return self.For.render( + index_name=self.visit(node.index_name, **kwargs), + iter_start=self.visit(node.iter_start, **kwargs), + iter_stop=self.visit(node.iter_stop, **kwargs), + iter_step=self.visit(node.iter_step, **kwargs), + body=body, + ) + def visit_VerticalPass(self, node: npir.VerticalPass, **kwargs): is_serial = node.direction != common.LoopOrder.PARALLEL has_variable_k = bool(node.walk_values().if_isinstance(npir.VarKOffset).to_list()) diff --git a/src/gt4py/cartesian/gtc/numpy/oir_to_npir.py b/src/gt4py/cartesian/gtc/numpy/oir_to_npir.py index b6aeb49823..9477882a80 100644 --- a/src/gt4py/cartesian/gtc/numpy/oir_to_npir.py +++ b/src/gt4py/cartesian/gtc/numpy/oir_to_npir.py @@ -78,6 +78,9 @@ def visit_VariableKOffset( ) -> Tuple[int, int, eve.Node]: return 0, 0, npir.VarKOffset(k=self.visit(node.k, **kwargs)) + def visit_ForIndex(self, node: oir.ForIndex, **kwargs: Any) -> npir.ForIndex: + return npir.ForIndex(name=node.name, dtype=node.dtype) + def visit_FieldAccess(self, node: oir.FieldAccess, **kwargs: Any) -> npir.FieldSlice: i_offset, j_offset, k_offset = self.visit(node.offset, **kwargs) data_index = [self.visit(index, **kwargs) for index in node.data_index] @@ -97,7 +100,9 @@ def visit_BinaryOp( self, node: oir.BinaryOp, **kwargs: Any ) -> Union[npir.VectorArithmetic, npir.VectorLogic]: args = dict( - op=node.op, left=self.visit(node.left, **kwargs), right=self.visit(node.right, **kwargs) + op=node.op, + left=self.visit(node.left, **kwargs), + right=self.visit(node.right, **kwargs), ) if isinstance(node.op, common.LogicalOperator): return npir.VectorLogic(**args) @@ -162,7 +167,17 @@ def visit_While( cond_expr = npir.VectorLogic(op=common.LogicalOperator.AND, left=mask, right=cond_expr) return npir.While( - cond=cond_expr, body=utils.flatten_list(self.visit(node.body, mask=cond_expr, **kwargs)) + cond=cond_expr, + body=utils.flatten_list(self.visit(node.body, mask=cond_expr, **kwargs)), + ) + + def visit_For(self, node: oir.For, **kwargs: Any) -> npir.For: + return npir.For( + index_name=node.index_name, + iter_start=node.iter_start, + iter_stop=node.iter_stop, + iter_step=node.iter_step, + body=utils.flatten_list(self.visit(node.body, **kwargs)), ) def visit_HorizontalRestriction( @@ -191,11 +206,17 @@ def visit_HorizontalExecution( stmts = utils.flatten_list(self.visit(node.body, extent=extent, **kwargs)) return npir.HorizontalBlock( - body=stmts, extent=extent, declarations=self.visit(node.declarations, **kwargs) + body=stmts, + extent=extent, + declarations=self.visit(node.declarations, **kwargs), ) def visit_VerticalLoopSection( - self, node: oir.VerticalLoopSection, *, loop_order: common.LoopOrder, **kwargs: Any + self, + node: oir.VerticalLoopSection, + *, + loop_order: common.LoopOrder, + **kwargs: Any, ) -> npir.VerticalPass: return npir.VerticalPass( body=self.visit(node.horizontal_executions, **kwargs), From cfa85f4a68d08d37fa2538a3da51a7b277386b5c Mon Sep 17 00:00:00 2001 From: Gabriel Vollenweider Date: Sat, 6 Jun 2026 13:30:11 +0200 Subject: [PATCH 33/48] avoid int overflow for arrays with many gridpoints --- src/gt4py/storage/allocators.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/storage/allocators.py b/src/gt4py/storage/allocators.py index 394374c2a4..5cdc904e54 100644 --- a/src/gt4py/storage/allocators.py +++ b/src/gt4py/storage/allocators.py @@ -212,7 +212,7 @@ def allocate( # Compute the padding required in the contiguous dimension to get aligned blocks dims_layout = [layout_map.index(i) for i in range(len(shape))] # Convert shape size to same data type (note that `np.int16` can overflow) - padded_shape_lst = [np.int32(x) for x in shape] + padded_shape_lst = [np.int64(x) for x in shape] if ndim > 0: padded_shape_lst[dims_layout[-1]] = ( # type: ignore[call-overload] math.ceil(shape[dims_layout[-1]] / items_per_aligned_block) From 41300dc91a44d3662f19862d03dba8a1683e0730 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 10 Sep 2026 13:48:26 +0200 Subject: [PATCH 34/48] Pass extra cxx compile args to cuda/rocm compiler --- src/gt4py/cartesian/backend/pyext_builder.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/src/gt4py/cartesian/backend/pyext_builder.py b/src/gt4py/cartesian/backend/pyext_builder.py index 6f91be5d46..7186dd05a4 100644 --- a/src/gt4py/cartesian/backend/pyext_builder.py +++ b/src/gt4py/cartesian/backend/pyext_builder.py @@ -98,6 +98,11 @@ def get_gt_pyext_build_opts( "-std=c++20", f"-ftemplate-depth={gt_config.build_settings['cpp_template_depth']}", *extra_compile_args_from_config["cuda"], + *( + arg + for cxx_compile_arg in extra_compile_args["cxx"] + for arg in f"--compiler-options {cxx_compile_arg}".split(" ") + ), ] if is_rocm_gpu: extra_compile_args["cuda"] += [ From e5805fd61b0f7827112eb6bcdae39f4d6ad84682 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Mon, 21 Sep 2026 09:35:34 +0200 Subject: [PATCH 35/48] Revert "feat[cartesian]: DaCe optimal for/map schedule (#2628)" This reverts commit 20d1bea8e20bf4d5bc004cb3222d97de7ec9f4e0. # Conflicts: # src/gt4py/cartesian/gtc/dace/oir_to_treeir.py # tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py --- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 48 ++++---- .../stencil_definitions.py | 2 +- .../test_code_generation.py | 107 ------------------ 3 files changed, 28 insertions(+), 129 deletions(-) diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index 4c0b2c78c8..46ad37bfb3 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -31,8 +31,10 @@ """Default dace residency types per device type.""" -def _resolve_map_schedule(device_type: dtypes.DeviceType) -> dtypes.ScheduleType: - """Optimal kernel schedule type based on target device.""" +def _resolve_default_map_schedule( + device_type: dtypes.DeviceType, +) -> dtypes.ScheduleType: + """Default kernel target per device type.""" if device_type == dtypes.DeviceType.GPU: return dtypes.ScheduleType.GPU_Device @@ -42,7 +44,7 @@ def _resolve_map_schedule(device_type: dtypes.DeviceType) -> dtypes.ScheduleType if not gt_config.build_settings["openmp"]["use_openmp"]: return dtypes.ScheduleType.Sequential - return dtypes.ScheduleType.CPU_Multicore + return dtypes.ScheduleType.Default class OIRToTreeIR(eve.NodeVisitor): @@ -146,7 +148,7 @@ def visit_HorizontalExecution(self, node: oir.HorizontalExecution, ctx: tir.Cont loop = tir.HorizontalLoop( bounds_i=tir.Bounds(start=axis_start_i, end=axis_end_i), bounds_j=tir.Bounds(start=axis_start_j, end=axis_end_j), - schedule=_resolve_map_schedule(self._device_type), + schedule=_resolve_default_map_schedule(self._device_type), children=[], parent=ctx.current_scope, ) @@ -281,6 +283,19 @@ def visit_Interval( return tir.Bounds(start=start, end=end) + def _vertical_loop_schedule(self) -> dtypes.ScheduleType: + """ + Defines the vertical loop schedule. + + Current strategy is to + - keep the vertical loop on the host for both, CPU and GPU targets + - and run it in parallel on CPU and sequential on GPU. + """ + if self._device_type == dtypes.DeviceType.GPU: + return dtypes.ScheduleType.Sequential + + return _resolve_default_map_schedule(self._device_type) + def visit_VerticalLoopSection( self, node: oir.VerticalLoopSection, ctx: tir.Context, loop_order: common.LoopOrder ) -> None: @@ -291,23 +306,14 @@ def visit_VerticalLoopSection( axis_end=tir.Axis.K.domain_dace_symbol(), ) - loop: tir.SequentialVerticalLoop | tir.ParallelVerticalLoop - if loop_order == common.LoopOrder.PARALLEL: - loop = tir.ParallelVerticalLoop( - iteration_variable=tir.Axis.K.iteration_symbol(), - bounds_k=bounds, - schedule=_resolve_map_schedule(self._device_type), - children=[], - parent=ctx.current_scope, - ) - else: - loop = tir.SequentialVerticalLoop( - iteration_variable=tir.Axis.K.iteration_symbol(), - bounds_k=bounds, - loop_order=loop_order, - children=[], - parent=ctx.current_scope, - ) + loop = tir.VerticalLoop( + iteration_variable=eve.SymbolRef(f"{tir.Axis.K.iteration_symbol()}_{id(node)}"), + loop_order=loop_order, + bounds_k=bounds, + schedule=self._vertical_loop_schedule(), + children=[], + parent=ctx.current_scope, + ) with loop.scope(ctx): self.visit(node.horizontal_executions, ctx=ctx) diff --git a/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py b/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py index b3d412cb79..3c97675fd5 100644 --- a/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py +++ b/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py @@ -78,7 +78,7 @@ def copy_stencil(field_a: Field3D, field_b: Field3D): @gtscript.function def a_gtscript_function(b): - return sqrt(abs(b[0, 0, 0])) + return sqrt(abs(b[0, 1, 0])) @register diff --git a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py index f16c92c18a..5519ff790c 100644 --- a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py +++ b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py @@ -24,17 +24,11 @@ J, K, IJ, - IJK, computation, horizontal, interval, region, sin, - tan, - isfinite, - isinf, - isnan, - sqrt, ) from gt4py.storage.cartesian import utils as storage_utils @@ -1806,104 +1800,3 @@ def test_set_2d_mask( test_set_2d_mask(output, input, mask_2d) assert (mask_2d == 0).all() - - -@pytest.mark.parametrize( - "backend", - [ - "debug", - pytest.param("dace:cpu", marks=[pytest.mark.uses_dace]), - pytest.param( - "dace:gpu", - marks=[ - pytest.mark.uses_dace, - pytest.mark.requires_gpu, - pytest.mark.xfail( - raises=SystemExit, - reason="DaCe issue: Missing `_gbar` symbol for global sync inside nested SDFG.", - ), - ], - ), - pytest.param("gt:gpu", marks=[pytest.mark.requires_gpu]), - ], -) -def test_offset_j_in_temporaries(backend: str) -> None: - @gtscript.function - def a_gtscript_function(b): - return sqrt(abs(b[0, 1, 0])) - - @gtscript.stencil(backend=backend) - def test_stencil_offset_j_in_temporaries( - field_in: Field[IJK, np.float64], # type: ignore - field_out: Field[IJK, np.float64], # type: ignore - ) -> None: - with computation(PARALLEL), interval(...): - abs_res = abs(field_in) - tan_res = tan(abs_res) - - # This is an offset in J on a temporary, it will - # require a global kernel sync in KJI for GPU - sqrt_res = a_gtscript_function(tan_res) - - field_out = ( - sqrt_res - if isfinite(sqrt_res) - else field_in - if isinf(sqrt_res) - else field_out - if isnan(sqrt_res) - else 0.0 - ) - - -@gtscript.enum -class MyEnum(IntEnum): - Zero = 0 - A = 10 - B = 20 - C = 30 - - -@pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_enum_runtime(backend): - - @gtscript.stencil(backend=backend) - def the_stencil(out_field: Field[int], order: MyEnum): # type: ignore - with computation(PARALLEL), interval(0, 1): - out_field = 32 - if order < MyEnum.A: - out_field = MyEnum.A - - with computation(PARALLEL), interval(1, 2): - out_field = 23 - out_field = MyEnum.B - - with computation(PARALLEL), interval(2, None): - out_field = 56 - out_field = MyEnum.C - - domain = (5, 5, 5) - out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=int) - - the_stencil(out_arr, MyEnum.Zero) - - assert out_arr[0, 0, 0] == MyEnum.A.value - assert out_arr[0, 0, 1] == MyEnum.B.value - assert (out_arr[0, 0, 2:] == MyEnum.C.value).all() - - -@pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_negated_bool_runtime(backend): - - @gtscript.stencil(backend=backend) - def the_stencil(out_field: Field[int], done: bool): # type: ignore - with computation(PARALLEL), interval(...): - if not done: - out_field = 1 - - domain = (5, 5, 5) - out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=int) - - the_stencil(out_arr, done=False) - - assert (out_arr[:] == 1).all() From b04e352c3572a870ee087e2e628f7d2ca6f0c5dd Mon Sep 17 00:00:00 2001 From: stubbiali Date: Mon, 21 Sep 2026 10:08:54 +0200 Subject: [PATCH 36/48] Revert "refactor[cartesian]: update treeir representation of the vertical loop (#2688)" This reverts commit bd1e156d400df85a47190ed2142aa737c68c64b9. # Conflicts: # src/gt4py/cartesian/gtc/dace/oir_to_treeir.py # src/gt4py/cartesian/gtc/dace/treeir.py # tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py --- src/gt4py/cartesian/gtc/dace/treeir.py | 13 +- .../cartesian/gtc/dace/treeir_to_stree.py | 18 +- .../test_code_generation.py | 288 ++++++++---------- 3 files changed, 148 insertions(+), 171 deletions(-) diff --git a/src/gt4py/cartesian/gtc/dace/treeir.py b/src/gt4py/cartesian/gtc/dace/treeir.py index b838c0bc25..d22052e13e 100644 --- a/src/gt4py/cartesian/gtc/dace/treeir.py +++ b/src/gt4py/cartesian/gtc/dace/treeir.py @@ -126,15 +126,16 @@ class HorizontalLoop(TreeScope): schedule: dtypes.ScheduleType -class SequentialVerticalLoop(TreeScope): +class VerticalLoop(TreeScope): iteration_variable: eve.SymbolRef - bounds_k: Bounds + """ + DaCe 1.x (without CFGs) maps sequential loops to a state machine with the iteration variable + on interstate edges. Having unique symbols makes DaCe 1.x happy and allows to rename symbols + via search & replace. + """ loop_order: common.LoopOrder - - -class ParallelVerticalLoop(TreeScope): - iteration_variable: eve.SymbolRef bounds_k: Bounds + schedule: dtypes.ScheduleType diff --git a/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py b/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py index 47df922ac6..260f5fe971 100644 --- a/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py +++ b/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py @@ -84,13 +84,17 @@ def visit_HorizontalLoop(self, node: tir.HorizontalLoop, ctx: Context) -> None: with ContextPushPop(ctx, map_scope): self.visit(node.children, ctx=ctx) - def visit_SequentialVerticalLoop(self, node: tir.SequentialVerticalLoop, ctx: Context) -> None: - for_scope = tn.ForScope(loop=_loop_region_for(node, ctx), children=[]) + def visit_VerticalLoop(self, node: tir.VerticalLoop, ctx: Context) -> None: + # For serial loops, create a ForScope and add it to the tree + if node.loop_order != common.LoopOrder.PARALLEL: + for_scope = tn.ForScope(loop=_loop_region_for(node, ctx), children=[]) - with ContextPushPop(ctx, for_scope): - self.visit(node.children, ctx=ctx) + with ContextPushPop(ctx, for_scope): + self.visit(node.children, ctx=ctx) + + return - def visit_ParallelVerticalLoop(self, node: tir.ParallelVerticalLoop, ctx: Context) -> None: + # For parallel loops, create a map and add it to the tree dace_map = nodes.Map( label=f"{ctx.tree.name}__v_map_{id(node)}", params=[node.iteration_variable], @@ -136,9 +140,9 @@ def visit_TreeRoot(self, node: tir.TreeRoot) -> tn.ScheduleTreeRoot: return ctx.tree -def _loop_region_for(node: tir.SequentialVerticalLoop, ctx: Context) -> LoopRegion: +def _loop_region_for(node: tir.VerticalLoop, ctx: Context) -> LoopRegion: """ - Translates a sequential vertical loop into a Dace LoopRegion to be used in `tn.ForScope`. + Translates a vertical loop into a Dace LoopRegion to be used in `tn.ForScope`. :param node: Vertical loop to translate :return: DaCe LoopRegion to use in `tn.ForScope` diff --git a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py index 5519ff790c..f7a00fddc0 100644 --- a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py +++ b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py @@ -44,7 +44,7 @@ @pytest.mark.parametrize("name", stencil_definitions) @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_generation(name, backend) -> None: +def test_generation(name, backend): stencil_definition = stencil_definitions[name] externals = externals_registry[name] stencil = gtscript.stencil(backend, stencil_definition, externals=externals) @@ -65,17 +65,17 @@ def test_generation(name, backend) -> None: @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lazy_stencil(backend) -> None: +def test_lazy_stencil(backend): @gtscript.lazy_stencil(backend=backend) - def definition(field_a: Field[np.float64], field_b: Field[np.float64]) -> None: # type: ignore + def definition(field_a: Field[np.float64], field_b: Field[np.float64]): # type: ignore with computation(PARALLEL), interval(...): field_a[0, 0, 0] = field_b @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_temporary_field_declared_in_if(backend) -> None: +def test_temporary_field_declared_in_if(backend): @gtscript.stencil(backend=backend) - def definition(field_a: Field[np.float64]) -> None: # type: ignore + def definition(field_a: Field[np.float64]): # type: ignore with computation(PARALLEL), interval(...): if field_a < 0: field_b = -field_a @@ -85,21 +85,21 @@ def definition(field_a: Field[np.float64]) -> None: # type: ignore @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_stage_without_effect(backend) -> None: +def test_stage_without_effect(backend): @gtscript.stencil(backend=backend) - def definition(field_a: Field[np.float64]) -> None: # type: ignore + def definition(field_a: Field[np.float64]): # type: ignore with computation(PARALLEL), interval(...): field_c = 0.0 # noqa: F841 -def test_ignore_np_errstate() -> None: - def setup_and_run(backend, **kwargs) -> None: +def test_ignore_np_errstate(): + def setup_and_run(backend, **kwargs): field_a = gt_storage.zeros( dtype=np.float64, backend=backend, shape=(3, 3, 1), aligned_index=(0, 0, 0) ) @gtscript.stencil(backend=backend, **kwargs) - def divide_by_zero(field_a: Field[np.float64]) -> None: # type: ignore + def divide_by_zero(field_a: Field[np.float64]): # type: ignore with computation(PARALLEL), interval(...): field_a = 1.0 / field_a @@ -113,12 +113,12 @@ def divide_by_zero(field_a: Field[np.float64]) -> None: # type: ignore @pytest.mark.parametrize("backend", CPU_BACKENDS) -def test_stencil_without_effect(backend) -> None: - def definition1(field_in: Field[np.float64]) -> None: # type: ignore +def test_stencil_without_effect(backend): + def definition1(field_in: Field[np.float64]): # type: ignore with computation(PARALLEL), interval(...): tmp = 0.0 # noqa: F841 - def definition2(f_in: Field[np.float64]) -> None: # type: ignore + def definition2(f_in: Field[np.float64]): # type: ignore from __externals__ import flag # type: ignore with computation(PARALLEL), interval(...): @@ -141,7 +141,7 @@ def definition2(f_in: Field[np.float64]) -> None: # type: ignore @pytest.mark.parametrize("backend", CPU_BACKENDS) -def test_stage_merger_induced_interval_block_reordering(backend) -> None: +def test_stage_merger_induced_interval_block_reordering(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(23, 23, 23), aligned_index=(0, 0, 0) ) @@ -150,7 +150,7 @@ def test_stage_merger_induced_interval_block_reordering(backend) -> None: ) @gtscript.stencil(backend=backend) - def stencil(field_in: Field[np.float64], field_out: Field[np.float64]) -> None: # type: ignore + def stencil(field_in: Field[np.float64], field_out: Field[np.float64]): # type: ignore with computation(BACKWARD): with interval(-2, -1): # block 1 field_out = field_in @@ -169,13 +169,13 @@ def stencil(field_in: Field[np.float64], field_out: Field[np.float64]) -> None: @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_inputs(backend) -> None: +def test_lower_dimensional_inputs(backend): @gtscript.stencil(backend=backend) def stencil( field_3d: Field[gtscript.IJK, np.float64], # type: ignore field_2d: Field[gtscript.IJ, np.float64], # type: ignore field_1d: Field[gtscript.K, np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL): with interval(0, -1): tmp = field_2d + field_1d[1] @@ -223,13 +223,13 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_masked(backend) -> None: +def test_lower_dimensional_masked(backend): @gtscript.stencil(backend=backend) def copy_2to3( cond: Field[gtscript.IJK, np.float64], # type: ignore inp: Field[gtscript.IJ, np.float64], # type: ignore outp: Field[gtscript.IJK, np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): if cond > 0.0: outp[0, 0, 0] = inp @@ -254,13 +254,13 @@ def copy_2to3( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_masked_2dcond(backend) -> None: +def test_lower_dimensional_masked_2dcond(backend): @gtscript.stencil(backend=backend) def copy_2to3( cond: Field[gtscript.IJK, np.float64], # type: ignore inp: Field[gtscript.IJ, np.float64], # type: ignore outp: Field[gtscript.IJK, np.float64], # type: ignore - ) -> None: + ): with computation(FORWARD), interval(...): if cond > 0.0: outp[0, 0, 0] = inp @@ -285,12 +285,12 @@ def copy_2to3( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_inputs_2d_to_3d_forward(backend) -> None: +def test_lower_dimensional_inputs_2d_to_3d_forward(backend): @gtscript.stencil(backend=backend) def copy_2to3( inp: Field[gtscript.IJ, np.float64], # type: ignore outp: Field[gtscript.IJK, np.float64], # type: ignore - ) -> None: + ): with computation(FORWARD), interval(...): outp[0, 0, 0] = inp @@ -307,7 +307,7 @@ def copy_2to3( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_higher_dimensional_fields(backend) -> None: +def test_higher_dimensional_fields(backend): FLOAT64_VEC2 = (np.float64, (2,)) FLOAT64_MAT22 = (np.float64, (2, 2)) @@ -316,7 +316,7 @@ def stencil( field: Field[np.float64], # type: ignore vec_field: Field[FLOAT64_VEC2], # type: ignore mat_field: Field[FLOAT64_MAT22], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): tmp = vec_field[0, 0, 0][0] + vec_field[0, 0, 0][1] # noqa: F841 @@ -362,13 +362,13 @@ def stencil( @pytest.mark.parametrize("backend", CPU_BACKENDS) -def test_input_order(backend) -> None: +def test_input_order(backend): @gtscript.stencil(backend=backend) def stencil( in_field: Field[np.float64], # type: ignore parameter: np.float64, out_field: Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field * parameter @@ -385,13 +385,13 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_variable_offsets(backend) -> None: +def test_variable_offsets(backend): @gtscript.stencil(backend=backend) def stencil_ij( in_field: Field[np.float64], # type: ignore out_field: Field[np.float64], # type: ignore index_field: Field[gtscript.IJ, int], # type: ignore - ) -> None: + ): with computation(FORWARD), interval(...): out_field[0, 0, 0] = in_field[0, 0, 1] + in_field[0, 0, index_field + 1] index_field = index_field + 1 @@ -401,13 +401,13 @@ def stencil_ijk( in_field: Field[np.float64], # type: ignore out_field: Field[np.float64], # type: ignore index_field: Field[int], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, 1] + in_field[0, 0, index_field + 1] @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_variable_offsets_and_while_loop(backend) -> None: +def test_variable_offsets_and_while_loop(backend): @gtscript.stencil(backend=backend) def stencil( pe1: Field[np.float64], # type: ignore @@ -415,7 +415,7 @@ def stencil( qin: Field[np.float64], # type: ignore qout: Field[np.float64], # type: ignore lev: Field[gtscript.IJ, np.int_], # type: ignore - ) -> None: + ): with computation(FORWARD), interval(0, -1): if pe2[0, 0, 1] <= pe1[0, 0, lev]: qout = qin[0, 0, 1] @@ -428,7 +428,7 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_nested_while_loop(backend) -> None: +def test_nested_while_loop(backend): @gtscript.stencil(backend=backend) def stencil(field_a: Field[np.float64], field_b: Field[np.int_]): # type: ignore with computation(PARALLEL), interval(...): @@ -440,9 +440,9 @@ def stencil(field_a: Field[np.float64], field_b: Field[np.int_]): # type: ignor @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_mask_with_offset_written_in_conditional(backend) -> None: +def test_mask_with_offset_written_in_conditional(backend): @gtscript.stencil(backend) - def stencil(outp: Field[np.float64]) -> None: # type: ignore + def stencil(outp: Field[np.float64]): # type: ignore with computation(PARALLEL), interval(...): cond = True if cond[0, -1, 0] or cond[0, 0, 0]: @@ -461,14 +461,14 @@ def stencil(outp: Field[np.float64]) -> None: # type: ignore @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_write_data_dim_indirect_addressing(backend) -> None: +def test_write_data_dim_indirect_addressing(backend): INT32_VEC2 = (np.int32, (2,)) def stencil( input_field: Field[gtscript.IJK, np.int32], # type: ignore output_field: Field[gtscript.IJK, INT32_VEC2], # type: ignore index: int, - ) -> None: + ): with computation(PARALLEL), interval(...): output_field[0, 0, 0][index] = input_field @@ -486,14 +486,14 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_read_data_dim_indirect_addressing(backend) -> None: +def test_read_data_dim_indirect_addressing(backend): INT32_VEC2 = (np.int32, (2,)) def stencil( input_field: Field[gtscript.IJK, INT32_VEC2], # type: ignore output_field: Field[gtscript.IJK, np.int32], # type: ignore index: int, - ) -> None: + ): with computation(PARALLEL), interval(...): output_field[0, 0, 0] = input_field[0, 0, 0][index] @@ -512,12 +512,12 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) class TestNegativeOrigin: - def test_negative_origin_i(self, backend) -> None: + def test_negative_origin_i(self, backend): @gtscript.stencil(backend=backend) def stencil_i( input_field: Field[gtscript.IJK, np.int32], # type: ignore output_field: Field[gtscript.IJK, np.int32], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): output_field[0, 0, 0] = input_field[1, 0, 0] @@ -531,12 +531,12 @@ def stencil_i( stencil_i(input_field, output_field, origin={"input_field": (-1, 0, 0)}) assert output_field[0, 0, 0] == 1 - def test_negative_origin_k(self, backend) -> None: + def test_negative_origin_k(self, backend): @gtscript.stencil(backend=backend) def stencil_k( input_field: Field[gtscript.IJK, np.int32], # type: ignore output_field: Field[gtscript.IJK, np.int32], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): output_field[0, 0, 0] = input_field[0, 0, 1] @@ -552,9 +552,9 @@ def stencil_k( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_origin_k_fields(backend) -> None: +def test_origin_k_fields(backend): @gtscript.stencil(backend=backend, rebuild=True) - def k_to_ijk(outp: Field[np.float64], inp: Field[gtscript.K, np.float64]) -> None: # type: ignore + def k_to_ijk(outp: Field[np.float64], inp: Field[gtscript.K, np.float64]): # type: ignore with computation(PARALLEL), interval(...): outp[0, 0, 0] = inp @@ -580,7 +580,7 @@ def k_to_ijk(outp: Field[np.float64], inp: Field[gtscript.K, np.float64]) -> Non @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_tmp_stencil(backend) -> None: +def test_tmp_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(6, 6, 6), aligned_index=(0, 0, 0) ) @@ -589,10 +589,7 @@ def test_tmp_stencil(backend) -> None: ) @gtscript.stencil(backend=backend) - def stencil( - field_in: gtscript.Field[np.float64], # type: ignore - field_out: gtscript.Field[np.float64], # type: ignore - ) -> None: + def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.float64]): # type: ignore with computation(PARALLEL): with interval(...): tmp = field_in + 1 @@ -613,7 +610,7 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_backward_stencil(backend) -> None: +def test_backward_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -622,10 +619,7 @@ def test_backward_stencil(backend) -> None: ) @gtscript.stencil(backend=backend) - def stencil( - field_in: gtscript.Field[np.float64], # type: ignore - field_out: gtscript.Field[np.float64], # type: ignore - ) -> None: + def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.float64]): # type: ignore with computation(BACKWARD): with interval(-1, None): field_in = 2 @@ -644,7 +638,7 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_while_stencil(backend) -> None: +def test_while_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(6, 6, 6), aligned_index=(0, 0, 0) ) @@ -653,10 +647,7 @@ def test_while_stencil(backend) -> None: ) @gtscript.stencil(backend=backend) - def stencil( - field_in: gtscript.Field[np.float64], # type: ignore - field_out: gtscript.Field[np.float64], # type: ignore - ) -> None: + def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.float64]): # type: ignore with computation(PARALLEL): with interval(...): while field_in < 10: @@ -671,7 +662,7 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_higher_dim_literal_stencil(backend) -> None: +def test_higher_dim_literal_stencil(backend): FLOAT64_NDDIM = (np.float64, (4,)) field_in = gt_storage.ones( @@ -686,7 +677,7 @@ def test_higher_dim_literal_stencil(backend) -> None: def stencil( vec_field: gtscript.Field[FLOAT64_NDDIM], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): out_field[0, 0, 0] = vec_field[0, 0, 0][2] @@ -698,7 +689,7 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_higher_dim_scalar_stencil(backend) -> None: +def test_higher_dim_scalar_stencil(backend): FLOAT64_NDDIM = (np.float64, (4,)) field_in = gt_storage.ones( @@ -714,7 +705,7 @@ def stencil( vec_field: gtscript.Field[FLOAT64_NDDIM], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore scalar_argument: int, - ) -> None: + ): with computation(PARALLEL), interval(...): out_field[0, 0, 0] = vec_field[0, 0, 0][scalar_argument] @@ -744,7 +735,7 @@ def data_dims_with_numpy_int_type( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_native_function_call_stencil(backend) -> None: +def test_native_function_call_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -756,7 +747,7 @@ def test_native_function_call_stencil(backend) -> None: def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, 0] + sin(0.848062) @@ -766,7 +757,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_unary_operator_stencil(backend) -> None: +def test_unary_operator_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -778,7 +769,7 @@ def test_unary_operator_stencil(backend) -> None: def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): out_field[0, 0, 0] = -in_field[0, 0, 0] @@ -788,7 +779,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_ternary_operator_stencil(backend) -> None: +def test_ternary_operator_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -801,7 +792,7 @@ def test_ternary_operator_stencil(backend) -> None: def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, 0] if in_field > 10 else in_field[0, 0, 0] + 1 @@ -813,7 +804,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_mask_stencil(backend) -> None: +def test_mask_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -826,7 +817,7 @@ def test_mask_stencil(backend) -> None: def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): if in_field[0, 0, 0] > 0: out_field[0, 0, 0] = in_field @@ -840,7 +831,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_k_offset_stencil(backend) -> None: +def test_k_offset_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -855,7 +846,7 @@ def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore scalar_value: int, - ) -> None: + ): with computation(PARALLEL), interval(1, None): out_field[0, 0, 0] = in_field[0, 0, scalar_value] @@ -866,7 +857,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_k_offset_field_stencil(backend) -> None: +def test_k_offset_field_stencil(backend): field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -882,7 +873,7 @@ def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore idx_field: gtscript.Field[gtscript.IJ, np.int64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(1, None): out_field[0, 0, 0] = in_field[0, 0, idx_field + 1] @@ -893,7 +884,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_k_only_access_stencil(backend) -> None: +def test_k_only_access_stencil(backend): field_in = gt_storage.from_array( np.array([2, 3, 4, 5]), dtype=np.float64, backend=backend, aligned_index=(0,) ) @@ -905,7 +896,7 @@ def test_k_only_access_stencil(backend) -> None: def test_stencil( in_field: gtscript.Field[gtscript.K, np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL): with interval(0, 1): out_field[0, 0, 0] = in_field[1] @@ -919,7 +910,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_table_access_stencil(backend) -> None: +def test_table_access_stencil(backend): table_view = gt_storage.from_array( np.array([2, 3, 4, 5]), dtype=np.float64, backend=backend, aligned_index=(0,) ) @@ -931,7 +922,7 @@ def test_table_access_stencil(backend) -> None: def test_stencil( table_view: gtscript.GlobalTable[(np.float64, (4))], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL): with interval(0, 1): out_field[0, 0, 0] = table_view.A[1] @@ -945,9 +936,9 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_pruned_args_match(backend) -> None: +def test_pruned_args_match(backend): @gtscript.stencil(backend=backend) - def test(out: Field[np.float64], inp: Field[np.float64]) -> None: # type: ignore + def test(out: Field[np.float64], inp: Field[np.float64]): # type: ignore with computation(PARALLEL), interval(...): out = 0.0 with horizontal(region[I[0] - 1, J[0] - 1]): @@ -972,7 +963,7 @@ def test_K_offset_write_simple(backend: str) -> None: # A is untouched # B is written in K+1 and should have K_values, except for the first element (FORWARD) @gtscript.stencil(backend=backend) - def simple(A: Field[np.float64], B: Field[np.float64]) -> None: # type: ignore + def simple(A: Field[np.float64], B: Field[np.float64]): # type: ignore with computation(FORWARD), interval(...): B[0, 0, 1] = A @@ -1000,7 +991,7 @@ def test_K_offset_write_forward(backend: str) -> None: # means while A is update B will have non-updated values of A # Because of the interval, value of B[0] is 0 @gtscript.stencil(backend=backend) - def forward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64) -> None: # type: ignore + def forward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64): # type: ignore with computation(FORWARD), interval(1, None): A[0, 0, -1] = scalar B[0, 0, 0] = A @@ -1031,7 +1022,7 @@ def test_K_offset_write_backward(backend: str) -> None: # means A is update B will get the updated values of A # Because of the interval, B[0] is never written @gtscript.stencil(backend=backend) - def backward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64) -> None: # type: ignore + def backward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64): # type: ignore with computation(BACKWARD), interval(-1, None): A = scalar @@ -1055,17 +1046,13 @@ def backward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64) -> @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_K_offset_write_conditional(backend) -> None: +def test_K_offset_write_conditional(backend): arraylib = get_array_library(backend) array_shape = (1, 1, 4) K_values = arraylib.arange(start=40, stop=44) @gtscript.stencil(backend=backend) - def column_physics_conditional( - A: Field[np.float64], # type: ignore - B: Field[np.float64], # type: ignore - scalar: np.float64, - ) -> None: + def column_physics_conditional(A: Field[np.float64], B: Field[np.float64], scalar: np.float64): # type: ignore with computation(BACKWARD), interval(1, -1): if A > 0 and B > 0: A[0, 0, -1] = scalar @@ -1125,11 +1112,11 @@ def column_physics_conditional( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_direct_datadims_index(backend) -> None: +def test_direct_datadims_index(backend): F64_VEC4 = (np.float64, (2, 2, 2, 2)) @gtscript.stencil(backend=backend) - def test(out: Field[np.float64], inp: GlobalTable[F64_VEC4]) -> None: # type: ignore + def test(out: Field[np.float64], inp: GlobalTable[F64_VEC4]): # type: ignore with computation(PARALLEL), interval(...): out[0, 0, 0] = inp.A[1, 0, 1, 0] @@ -1141,7 +1128,7 @@ def test(out: Field[np.float64], inp: GlobalTable[F64_VEC4]) -> None: # type: i @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_function_inline_in_while(backend) -> None: +def test_function_inline_in_while(backend): @gtscript.function def add_42(v): return v + 42 @@ -1150,7 +1137,7 @@ def add_42(v): def test( in_field: Field[np.float64], # type: ignore out_field: Field[np.float64], # type: ignore - ) -> None: + ): with computation(PARALLEL), interval(...): count = 1 while count < 10: @@ -1166,21 +1153,21 @@ def test( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_cast_in_index(backend) -> None: +def test_cast_in_index(backend): @gtscript.stencil(backend) def cast_in_index( in_field: Field[np.float64], # type: ignore i32: np.int32, i64: np.int64, out_field: Field[np.float64], # type: ignore - ) -> None: + ): """Simple copy stencil with forced cast in index calculation.""" with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, i32 - i64] @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_read_after_write_stencil(backend) -> None: +def test_read_after_write_stencil(backend): """Stencil with multiple read after write access patterns.""" @gtscript.stencil(backend=backend) @@ -1194,7 +1181,7 @@ def lagrangian_contributions( q4_4: Field[np.float64], # type: ignore dp1: Field[np.float64], # type: ignore lev: Field[gtscript.IJ, np.int64], # type: ignore - ) -> None: + ): """ Args: q (out): @@ -1254,9 +1241,9 @@ def lagrangian_contributions( for backend in ["gt:cpu_ifirst", "numpy"] ], ) -def test_absolute_K_index_raise(backend) -> None: +def test_absolute_K_index_raise(backend): @gtscript.stencil(backend=backend) - def test_absolute_k_index(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore + def test_absolute_k_index(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=2) @@ -1269,7 +1256,7 @@ def test_absolute_k_index(in_field: Field[np.float64], out_field: Field[np.float pytest.param("dace:gpu", marks=[pytest.mark.uses_dace, pytest.mark.requires_gpu]), ], ) -def test_absolute_K_index(backend) -> None: +def test_absolute_K_index(backend): domain = (5, 5, 5) in_arr = gt_storage.ones(backend=backend, shape=domain, dtype=np.float64) @@ -1279,7 +1266,7 @@ def test_absolute_K_index(backend) -> None: out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=np.float64) @gtscript.stencil(backend=backend) - def test_literal_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore + def test_literal_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=2) @@ -1291,9 +1278,7 @@ def test_literal_access(in_field: Field[np.float64], out_field: Field[np.float64 @gtscript.stencil(backend=backend) def test_parameter_access( - in_field: Field[np.float64], # type:ignore - out_field: Field[np.float64], # type:ignore - idx: int, + in_field: Field[np.float64], out_field: Field[np.float64], idx: int ) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=idx) @@ -1305,9 +1290,9 @@ def test_parameter_access( assert (out_arr[:, :, :] == 42.42).all() @gtscript.stencil(backend=backend, externals={"K4": 4}) - def test_external_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore + def test_external_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: with computation(PARALLEL), interval(...): - from __externals__ import K4 # type:ignore + from __externals__ import K4 out_field = in_field.at(K=K4) @@ -1319,9 +1304,7 @@ def test_external_access(in_field: Field[np.float64], out_field: Field[np.float6 @gtscript.stencil(backend=backend) def test_field_access( - in_field: Field[np.float64], # type:ignore - index_field: Field[IJ, np.int64], # type:ignore - out_field: Field[np.float64], # type:ignore + in_field: Field[np.float64], index_field: Field[IJ, np.int64], out_field: Field[np.float64] ) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=index_field) @@ -1335,9 +1318,7 @@ def test_field_access( @gtscript.stencil(backend=backend) def test_field_access_computation( - in_field: Field[np.float64], # type:ignore - index_field: Field[IJ, np.int32], # type:ignore - out_field: Field[np.float64], # type:ignore + in_field: Field[np.float64], index_field: Field[IJ, np.int32], out_field: Field[np.float64] ) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=index_field - 1) @@ -1351,8 +1332,8 @@ def test_field_access_computation( @gtscript.stencil(backend=backend) def test_lower_dim_field( - k_field: Field[K, np.float64], # type:ignore - out_field: Field[np.float64], # type:ignore + k_field: Field[K, np.float64], + out_field: Field[np.float64], ) -> None: with computation(PARALLEL), interval(...): out_field = k_field.at(K=2) @@ -1365,8 +1346,8 @@ def test_lower_dim_field( @gtscript.stencil(backend=backend) def test_conditional_absolute( - in_field: Field[np.float64], # type:ignore - out_field: Field[np.float64], # type:ignore + in_field: Field[np.float64], + out_field: Field[np.float64], ) -> None: with computation(PARALLEL), interval(...): k_level = 0 @@ -1394,9 +1375,7 @@ def test_iterator_access(backend: str) -> None: @gtscript.stencil(backend=backend) def test_all_valid_usage( - field_A: Field[np.float64], # type:ignore - field_B: Field[np.float64], # type:ignore - offsets: Field[K, np.int32], # type:ignore + field_A: Field[np.float64], field_B: Field[np.float64], offsets: Field[K, np.int32] ) -> None: with computation(PARALLEL), interval(...): if K == 2: @@ -1429,7 +1408,7 @@ def test_iterator_access_raises_in_unsupported_backends(backend: str) -> None: field_B = gt_storage.zeros(backend=backend, shape=domain, dtype=np.float64) @gtscript.stencil(backend=backend) - def test_all_valid_usage(field_A: Field[np.float64], field_B: Field[np.float64]) -> None: # type:ignore + def test_all_valid_usage(field_A: Field[np.float64], field_B: Field[np.float64]) -> None: with computation(PARALLEL), interval(...): if K == 2: field_A = 20.20 @@ -1438,7 +1417,7 @@ def test_all_valid_usage(field_A: Field[np.float64], field_B: Field[np.float64]) test_all_valid_usage(field_A, field_B) -def test_runtime_interval_bounds() -> None: +def test_runtime_interval_bounds(): backend = "debug" domain = (5, 5, 10) in_arr = gt_storage.ones(backend=backend, shape=domain, dtype=np.float64) @@ -1543,14 +1522,14 @@ def test_temporary( ), ], ) -def test_runtime_interval_raises(backend) -> None: +def test_runtime_interval_raises(backend): @gtscript.stencil(backend=backend) def test_stencil( out_field: Field[np.float64], # type: ignore input_data: Field[np.float64], # type: ignore index_data: Field[gtscript.IJ, np.int64], # type: ignore scalar_arg: int, - ) -> None: + ): with computation(FORWARD), interval(0, 1): temporary: Field[IJ, np.float64] = 7 # type: ignore @@ -1573,16 +1552,16 @@ def test_stencil( pytest.param("dace:gpu", marks=[pytest.mark.uses_dace, pytest.mark.requires_gpu]), ], ) -def test_2d_temporaries(backend) -> None: +def test_2d_temporaries(backend): domain = (5, 5, 3) in_arr = gt_storage.ones(backend=backend, shape=domain, dtype=np.float64) out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=np.float64) @gtscript.stencil(backend=backend) - def test_with_plain_gt4py(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore + def test_with_plain_gt4py(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: with computation(FORWARD), interval(0, 1): - tmp_2D: Field[IJ, np.float64] = 0 # type:ignore + tmp_2D: Field[IJ, np.float64] = 0 with computation(FORWARD), interval(...): tmp_2D = tmp_2D + in_field @@ -1596,9 +1575,9 @@ def test_with_plain_gt4py(in_field: Field[np.float64], out_field: Field[np.float assert (out_arr[:, :, :] == domain[2]).all() @gtscript.stencil(backend=backend, dtypes={"MyFancySymbol": Field[IJ, np.float64]}) - def test_with_user_dtype(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore + def test_with_user_dtype(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: with computation(FORWARD), interval(0, 1): - tmp_2D: MyFancySymbol = 0 # type:ignore + tmp_2D: MyFancySymbol = 0 with computation(FORWARD), interval(...): out_field = tmp_2D @@ -1609,11 +1588,10 @@ def test_with_user_dtype(in_field: Field[np.float64], out_field: Field[np.float6 @gtscript.stencil(backend=backend) def test_failing_on_non_IJ( - in_field: Field[np.float64], # type:ignore - out_field: Field[np.float64], # type:ignore + in_field: Field[np.float64], out_field: Field[np.float64] ) -> None: with computation(FORWARD), interval(0, 1): - tmp_2D: Field[K, np.float64] = 0 # type:ignore + tmp_2D: Field[K, np.float64] = 0 with computation(FORWARD), interval(...): out_field = tmp_2D @@ -1631,11 +1609,11 @@ def test_failing_on_non_IJ( for backend in ["gt:cpu_ifirst", "gt:cpu_kfirst", "gt:gpu"] ], ) -def test_2d_temporaries_raises(backend) -> None: +def test_2d_temporaries_raises(backend): @gtscript.stencil(backend=backend) - def test_with_user_dtypes(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore + def test_with_user_dtypes(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: with computation(FORWARD), interval(0, 1): - tmp_2D: Field[IJ, np.float64] = 0 # type:ignore + tmp_2D: Field[IJ, np.float64] = 0 with computation(FORWARD), interval(...): out_field = tmp_2D @@ -1651,9 +1629,7 @@ def test_upcasting_both_sides_of_assignment(backend: str) -> None: @gtscript.stencil(backend=backend) def test_upcasting_stencil( - in_field: Field[np.float64], # type:ignore - index_field: Field[IJ, np.int32], # type:ignore - out_field: Field[np.float64], # type:ignore + in_field: Field[np.float64], index_field: Field[IJ, np.int32], out_field: Field[np.float64] ) -> None: with computation(FORWARD), interval(...): out_field[0, 0, index_field - 1] = in_field @@ -1662,7 +1638,7 @@ def test_upcasting_stencil( assert (input == output).all() -@pytest.mark.parametrize("backend", ALL_BACKENDS) +@pytest.mark.parametrize("backend", ("debug",)) # ALL_BACKENDS) def test_upcasting_leave_integer_power_arguments_alone(backend: str) -> None: domain = (5, 5, 5) @@ -1674,9 +1650,7 @@ def test_upcasting_leave_integer_power_arguments_alone(backend: str) -> None: @gtscript.stencil(backend=backend) def test_upcasting_stencil( - in_field: Field[np.float32], # type:ignore - squared: Field[IJ, np.int32], # type:ignore - out_field: Field[np.float32], # type:ignore + in_field: Field[np.float32], squared: Field[IJ, np.int32], out_field: Field[np.float32] ) -> None: with computation(FORWARD), interval(...): out_field = in_field**squared @@ -1688,14 +1662,14 @@ def test_no_write_and_read_with_horizontal_offset() -> None: with pytest.raises(ValueError, match="Self-assignment with offset in I or J is illegal."): @gtscript.stencil(backend="debug") - def self_assign_offset(field: Field[np.float64]) -> None: # type:ignore + def self_assign_offset(field: Field[np.float64]) -> None: with computation(PARALLEL), interval(...): field = (field[I - 1] + field[I + 1]) / 2 with pytest.raises(ValueError, match="Illegal write and read with horizontal offset"): @gtscript.stencil(backend="debug") - def self_assign_offset(field: Field[np.float64]) -> None: # type:ignore + def self_assign_offset(field: Field[np.float64]) -> None: with computation(PARALLEL), interval(...): tmp = (field[J - 1] + field[J + 1]) / 2 field = tmp * 2 @@ -1705,14 +1679,14 @@ def test_k_offsets_in_parallel_loops() -> None: with pytest.raises(ValueError, match="write and read with k-offsets in PARALLEL"): @gtscript.stencil(backend="debug") - def self_assign_offset_parallel(field: Field[np.int32]) -> None: # type:ignore + def self_assign_offset_parallel(field: Field[np.int32]) -> None: with computation(PARALLEL), interval(1, None): field = field[K - 1] * 2 with pytest.raises(ValueError, match="write and read with k-offsets in PARALLEL"): @gtscript.stencil(backend="debug") - def self_assign_offset_parallel_temp(field: Field[np.int32]) -> None: # type:ignore + def self_assign_offset_parallel_temp(field: Field[np.int32]) -> None: with computation(PARALLEL), interval(1, None): tmp = field[K - 1] field = tmp * 2 @@ -1722,7 +1696,7 @@ def self_assign_offset_parallel_temp(field: Field[np.int32]) -> None: # type:ig ): @gtscript.stencil(backend="debug") - def mixed_read_write(field: Field[np.int32]) -> None: # type:ignore + def mixed_read_write(field: Field[np.int32]): with computation(PARALLEL), interval(...): level = field.at(K=1) field = 2 * level @@ -1732,31 +1706,31 @@ def mixed_read_write(field: Field[np.int32]) -> None: # type:ignore ): @gtscript.stencil(backend="debug") - def mixed_read_write(field: Field[np.int32], offset: int = -1) -> None: # type:ignore + def mixed_read_write(field: Field[np.int32], offset: int = -1): with computation(PARALLEL), interval(1, None): bottom = field[0, 0, offset] field = field + 2 * bottom # center reads and writes are allowed @gtscript.stencil(backend="debug") - def self_assignment_center_read_parallel(field: Field[np.int32]) -> None: # type:ignore + def self_assignment_center_read_parallel(field: Field[np.int32]) -> None: with computation(PARALLEL), interval(...): field = field[0, 0, 0] * 2 @gtscript.stencil(backend="debug") - def self_assignment_center_write_parallel(field: Field[np.int32]) -> None: # type:ignore + def self_assignment_center_write_parallel(field: Field[np.int32]) -> None: with computation(PARALLEL), interval(...): field[0, 0, 0] = field * 2 # not mixing reads and writes are allowed (e.g. index fields) @gtscript.stencil(backend="debug") - def self_assignment_center_parallel(field: Field[np.float32], index: Field[np.int32]) -> None: # type:ignore + def self_assignment_center_parallel(field: Field[np.float32], index: Field[np.int32]) -> None: with computation(PARALLEL), interval(1, None): field = index + index[K - 1] * 2 # parallel intervals of static size 1 are allowed @gtscript.stencil(backend="debug") - def the_stencil(field: Field[np.bool_]) -> None: # type:ignore + def the_stencil(field: Field[np.bool_]) -> None: with computation(PARALLEL): with interval(0, 1): field = field[K + 1] @@ -1767,12 +1741,12 @@ def the_stencil(field: Field[np.bool_]) -> None: # type:ignore @pytest.mark.parametrize("backend", ALL_BACKENDS) def test_self_assignment_in_forward(backend: str) -> None: @gtscript.stencil(backend=backend) - def self_assignment_parallel(field: Field[np.int32]) -> None: # type:ignore + def self_assignment_parallel(field: Field[np.int32]) -> None: with computation(FORWARD), interval(1, None): field = field[K - 1] * 2 @gtscript.stencil(backend=backend) - def self_assignment_2_parallel(field: Field[np.int32]) -> None: # type:ignore + def self_assignment_2_parallel(field: Field[np.int32]) -> None: with computation(FORWARD), interval(1, None): tmp = field[K - 1] field = tmp * 2 @@ -1788,9 +1762,7 @@ def test_reset_mask_2d(backend: str) -> None: @gtscript.stencil(backend=backend) def test_set_2d_mask( - dp1: Field[np.float64], # type:ignore - pe1: Field[np.float64], # type:ignore - lev: Field[IJ, np.int32], # type:ignore + dp1: Field[np.float64], pe1: Field[np.float64], lev: Field[IJ, np.int32] ) -> None: with computation(PARALLEL), interval(0, -1): dp1 = pe1[0, 0, 1] - pe1 From 62b0896d85732c334652f4fa73d3dfd49774a77c Mon Sep 17 00:00:00 2001 From: stubbiali Date: Mon, 21 Sep 2026 10:19:00 +0200 Subject: [PATCH 37/48] Add gt4py and dace version to stencil fingerprint --- src/gt4py/cartesian/caching.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/src/gt4py/cartesian/caching.py b/src/gt4py/cartesian/caching.py index e3bdb280f3..62f950dc39 100644 --- a/src/gt4py/cartesian/caching.py +++ b/src/gt4py/cartesian/caching.py @@ -19,7 +19,9 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional from cached_property import cached_property +from dace import __version__ as dace_version +from gt4py import __version__ as gt4py_version from gt4py.cartesian import config as gt_config, utils as gt_utils from gt4py.cartesian.definitions import StencilID @@ -299,6 +301,7 @@ def _extract_api_annotations(self) -> List[str]: @property def stencil_id(self) -> StencilID: fingerprint = { + "gt4py_version": gt4py_version, "__main__": self.builder.definition._gtscript_["canonical_ast"], "docstring": inspect.getdoc(self.builder.definition), "api_annotations": f"[{', '.join(self._extract_api_annotations())}]", @@ -316,8 +319,10 @@ def stencil_id(self) -> StencilID: fingerprint["extra_compile_args"] = self.builder.options.backend_opts.get( "extra_compile_args", gt_config.GT4PY_EXTRA_COMPILE_ARGS ) - if self.builder.backend.name == "dace:gpu": - fingerprint["default_block_size"] = gt_config.DACE_DEFAULT_BLOCK_SIZE + if "dace" in self.builder.backend.name: + fingerprint["dace_version"] = dace_version + if self.builder.backend.name == "dace:gpu": + fingerprint["default_block_size"] = gt_config.DACE_DEFAULT_BLOCK_SIZE # ignore type because attrclass StencilID has generated constructor return StencilID( # type: ignore From 06160a46b4bb252f8e9caf3d377621c9d3e5cb56 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Mon, 21 Sep 2026 11:45:26 +0200 Subject: [PATCH 38/48] Fix vertical loop index symbol --- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index 46ad37bfb3..8d94b5e093 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -307,7 +307,7 @@ def visit_VerticalLoopSection( ) loop = tir.VerticalLoop( - iteration_variable=eve.SymbolRef(f"{tir.Axis.K.iteration_symbol()}_{id(node)}"), + iteration_variable=tir.Axis.K.iteration_symbol(), loop_order=loop_order, bounds_k=bounds, schedule=self._vertical_loop_schedule(), From 25f7d66b41b3d3929176682b73de5cd4b5fda20e Mon Sep 17 00:00:00 2001 From: stubbiali Date: Mon, 21 Sep 2026 17:29:22 +0200 Subject: [PATCH 39/48] Try moving vertical loop inside the kernel --- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 15 +-------------- 1 file changed, 1 insertion(+), 14 deletions(-) diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index 8d94b5e093..0f4a9d9b7c 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -283,19 +283,6 @@ def visit_Interval( return tir.Bounds(start=start, end=end) - def _vertical_loop_schedule(self) -> dtypes.ScheduleType: - """ - Defines the vertical loop schedule. - - Current strategy is to - - keep the vertical loop on the host for both, CPU and GPU targets - - and run it in parallel on CPU and sequential on GPU. - """ - if self._device_type == dtypes.DeviceType.GPU: - return dtypes.ScheduleType.Sequential - - return _resolve_default_map_schedule(self._device_type) - def visit_VerticalLoopSection( self, node: oir.VerticalLoopSection, ctx: tir.Context, loop_order: common.LoopOrder ) -> None: @@ -310,7 +297,7 @@ def visit_VerticalLoopSection( iteration_variable=tir.Axis.K.iteration_symbol(), loop_order=loop_order, bounds_k=bounds, - schedule=self._vertical_loop_schedule(), + schedule=_resolve_default_map_schedule(self._device_type), children=[], parent=ctx.current_scope, ) From 29d420153812da353d000741e26410d000543199 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 23 Sep 2026 10:35:01 +0200 Subject: [PATCH 40/48] Revert "Try moving vertical loop inside the kernel" This reverts commit 25f7d66b41b3d3929176682b73de5cd4b5fda20e. --- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index 0f4a9d9b7c..8d94b5e093 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -283,6 +283,19 @@ def visit_Interval( return tir.Bounds(start=start, end=end) + def _vertical_loop_schedule(self) -> dtypes.ScheduleType: + """ + Defines the vertical loop schedule. + + Current strategy is to + - keep the vertical loop on the host for both, CPU and GPU targets + - and run it in parallel on CPU and sequential on GPU. + """ + if self._device_type == dtypes.DeviceType.GPU: + return dtypes.ScheduleType.Sequential + + return _resolve_default_map_schedule(self._device_type) + def visit_VerticalLoopSection( self, node: oir.VerticalLoopSection, ctx: tir.Context, loop_order: common.LoopOrder ) -> None: @@ -297,7 +310,7 @@ def visit_VerticalLoopSection( iteration_variable=tir.Axis.K.iteration_symbol(), loop_order=loop_order, bounds_k=bounds, - schedule=_resolve_default_map_schedule(self._device_type), + schedule=self._vertical_loop_schedule(), children=[], parent=ctx.current_scope, ) From bd91ef7062ef2fb791a2546368a04e12483199a7 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 23 Sep 2026 10:36:33 +0200 Subject: [PATCH 41/48] Reapply "refactor[cartesian]: update treeir representation of the vertical loop (#2688)" This reverts commit b04e352c3572a870ee087e2e628f7d2ca6f0c5dd. --- src/gt4py/cartesian/gtc/dace/treeir.py | 13 +- .../cartesian/gtc/dace/treeir_to_stree.py | 18 +- .../test_code_generation.py | 288 ++++++++++-------- 3 files changed, 171 insertions(+), 148 deletions(-) diff --git a/src/gt4py/cartesian/gtc/dace/treeir.py b/src/gt4py/cartesian/gtc/dace/treeir.py index d22052e13e..b838c0bc25 100644 --- a/src/gt4py/cartesian/gtc/dace/treeir.py +++ b/src/gt4py/cartesian/gtc/dace/treeir.py @@ -126,16 +126,15 @@ class HorizontalLoop(TreeScope): schedule: dtypes.ScheduleType -class VerticalLoop(TreeScope): +class SequentialVerticalLoop(TreeScope): iteration_variable: eve.SymbolRef - """ - DaCe 1.x (without CFGs) maps sequential loops to a state machine with the iteration variable - on interstate edges. Having unique symbols makes DaCe 1.x happy and allows to rename symbols - via search & replace. - """ - loop_order: common.LoopOrder bounds_k: Bounds + loop_order: common.LoopOrder + +class ParallelVerticalLoop(TreeScope): + iteration_variable: eve.SymbolRef + bounds_k: Bounds schedule: dtypes.ScheduleType diff --git a/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py b/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py index 260f5fe971..47df922ac6 100644 --- a/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py +++ b/src/gt4py/cartesian/gtc/dace/treeir_to_stree.py @@ -84,17 +84,13 @@ def visit_HorizontalLoop(self, node: tir.HorizontalLoop, ctx: Context) -> None: with ContextPushPop(ctx, map_scope): self.visit(node.children, ctx=ctx) - def visit_VerticalLoop(self, node: tir.VerticalLoop, ctx: Context) -> None: - # For serial loops, create a ForScope and add it to the tree - if node.loop_order != common.LoopOrder.PARALLEL: - for_scope = tn.ForScope(loop=_loop_region_for(node, ctx), children=[]) + def visit_SequentialVerticalLoop(self, node: tir.SequentialVerticalLoop, ctx: Context) -> None: + for_scope = tn.ForScope(loop=_loop_region_for(node, ctx), children=[]) - with ContextPushPop(ctx, for_scope): - self.visit(node.children, ctx=ctx) - - return + with ContextPushPop(ctx, for_scope): + self.visit(node.children, ctx=ctx) - # For parallel loops, create a map and add it to the tree + def visit_ParallelVerticalLoop(self, node: tir.ParallelVerticalLoop, ctx: Context) -> None: dace_map = nodes.Map( label=f"{ctx.tree.name}__v_map_{id(node)}", params=[node.iteration_variable], @@ -140,9 +136,9 @@ def visit_TreeRoot(self, node: tir.TreeRoot) -> tn.ScheduleTreeRoot: return ctx.tree -def _loop_region_for(node: tir.VerticalLoop, ctx: Context) -> LoopRegion: +def _loop_region_for(node: tir.SequentialVerticalLoop, ctx: Context) -> LoopRegion: """ - Translates a vertical loop into a Dace LoopRegion to be used in `tn.ForScope`. + Translates a sequential vertical loop into a Dace LoopRegion to be used in `tn.ForScope`. :param node: Vertical loop to translate :return: DaCe LoopRegion to use in `tn.ForScope` diff --git a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py index f7a00fddc0..5519ff790c 100644 --- a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py +++ b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py @@ -44,7 +44,7 @@ @pytest.mark.parametrize("name", stencil_definitions) @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_generation(name, backend): +def test_generation(name, backend) -> None: stencil_definition = stencil_definitions[name] externals = externals_registry[name] stencil = gtscript.stencil(backend, stencil_definition, externals=externals) @@ -65,17 +65,17 @@ def test_generation(name, backend): @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lazy_stencil(backend): +def test_lazy_stencil(backend) -> None: @gtscript.lazy_stencil(backend=backend) - def definition(field_a: Field[np.float64], field_b: Field[np.float64]): # type: ignore + def definition(field_a: Field[np.float64], field_b: Field[np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): field_a[0, 0, 0] = field_b @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_temporary_field_declared_in_if(backend): +def test_temporary_field_declared_in_if(backend) -> None: @gtscript.stencil(backend=backend) - def definition(field_a: Field[np.float64]): # type: ignore + def definition(field_a: Field[np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): if field_a < 0: field_b = -field_a @@ -85,21 +85,21 @@ def definition(field_a: Field[np.float64]): # type: ignore @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_stage_without_effect(backend): +def test_stage_without_effect(backend) -> None: @gtscript.stencil(backend=backend) - def definition(field_a: Field[np.float64]): # type: ignore + def definition(field_a: Field[np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): field_c = 0.0 # noqa: F841 -def test_ignore_np_errstate(): - def setup_and_run(backend, **kwargs): +def test_ignore_np_errstate() -> None: + def setup_and_run(backend, **kwargs) -> None: field_a = gt_storage.zeros( dtype=np.float64, backend=backend, shape=(3, 3, 1), aligned_index=(0, 0, 0) ) @gtscript.stencil(backend=backend, **kwargs) - def divide_by_zero(field_a: Field[np.float64]): # type: ignore + def divide_by_zero(field_a: Field[np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): field_a = 1.0 / field_a @@ -113,12 +113,12 @@ def divide_by_zero(field_a: Field[np.float64]): # type: ignore @pytest.mark.parametrize("backend", CPU_BACKENDS) -def test_stencil_without_effect(backend): - def definition1(field_in: Field[np.float64]): # type: ignore +def test_stencil_without_effect(backend) -> None: + def definition1(field_in: Field[np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): tmp = 0.0 # noqa: F841 - def definition2(f_in: Field[np.float64]): # type: ignore + def definition2(f_in: Field[np.float64]) -> None: # type: ignore from __externals__ import flag # type: ignore with computation(PARALLEL), interval(...): @@ -141,7 +141,7 @@ def definition2(f_in: Field[np.float64]): # type: ignore @pytest.mark.parametrize("backend", CPU_BACKENDS) -def test_stage_merger_induced_interval_block_reordering(backend): +def test_stage_merger_induced_interval_block_reordering(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(23, 23, 23), aligned_index=(0, 0, 0) ) @@ -150,7 +150,7 @@ def test_stage_merger_induced_interval_block_reordering(backend): ) @gtscript.stencil(backend=backend) - def stencil(field_in: Field[np.float64], field_out: Field[np.float64]): # type: ignore + def stencil(field_in: Field[np.float64], field_out: Field[np.float64]) -> None: # type: ignore with computation(BACKWARD): with interval(-2, -1): # block 1 field_out = field_in @@ -169,13 +169,13 @@ def stencil(field_in: Field[np.float64], field_out: Field[np.float64]): # type: @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_inputs(backend): +def test_lower_dimensional_inputs(backend) -> None: @gtscript.stencil(backend=backend) def stencil( field_3d: Field[gtscript.IJK, np.float64], # type: ignore field_2d: Field[gtscript.IJ, np.float64], # type: ignore field_1d: Field[gtscript.K, np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL): with interval(0, -1): tmp = field_2d + field_1d[1] @@ -223,13 +223,13 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_masked(backend): +def test_lower_dimensional_masked(backend) -> None: @gtscript.stencil(backend=backend) def copy_2to3( cond: Field[gtscript.IJK, np.float64], # type: ignore inp: Field[gtscript.IJ, np.float64], # type: ignore outp: Field[gtscript.IJK, np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): if cond > 0.0: outp[0, 0, 0] = inp @@ -254,13 +254,13 @@ def copy_2to3( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_masked_2dcond(backend): +def test_lower_dimensional_masked_2dcond(backend) -> None: @gtscript.stencil(backend=backend) def copy_2to3( cond: Field[gtscript.IJK, np.float64], # type: ignore inp: Field[gtscript.IJ, np.float64], # type: ignore outp: Field[gtscript.IJK, np.float64], # type: ignore - ): + ) -> None: with computation(FORWARD), interval(...): if cond > 0.0: outp[0, 0, 0] = inp @@ -285,12 +285,12 @@ def copy_2to3( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_lower_dimensional_inputs_2d_to_3d_forward(backend): +def test_lower_dimensional_inputs_2d_to_3d_forward(backend) -> None: @gtscript.stencil(backend=backend) def copy_2to3( inp: Field[gtscript.IJ, np.float64], # type: ignore outp: Field[gtscript.IJK, np.float64], # type: ignore - ): + ) -> None: with computation(FORWARD), interval(...): outp[0, 0, 0] = inp @@ -307,7 +307,7 @@ def copy_2to3( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_higher_dimensional_fields(backend): +def test_higher_dimensional_fields(backend) -> None: FLOAT64_VEC2 = (np.float64, (2,)) FLOAT64_MAT22 = (np.float64, (2, 2)) @@ -316,7 +316,7 @@ def stencil( field: Field[np.float64], # type: ignore vec_field: Field[FLOAT64_VEC2], # type: ignore mat_field: Field[FLOAT64_MAT22], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): tmp = vec_field[0, 0, 0][0] + vec_field[0, 0, 0][1] # noqa: F841 @@ -362,13 +362,13 @@ def stencil( @pytest.mark.parametrize("backend", CPU_BACKENDS) -def test_input_order(backend): +def test_input_order(backend) -> None: @gtscript.stencil(backend=backend) def stencil( in_field: Field[np.float64], # type: ignore parameter: np.float64, out_field: Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field * parameter @@ -385,13 +385,13 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_variable_offsets(backend): +def test_variable_offsets(backend) -> None: @gtscript.stencil(backend=backend) def stencil_ij( in_field: Field[np.float64], # type: ignore out_field: Field[np.float64], # type: ignore index_field: Field[gtscript.IJ, int], # type: ignore - ): + ) -> None: with computation(FORWARD), interval(...): out_field[0, 0, 0] = in_field[0, 0, 1] + in_field[0, 0, index_field + 1] index_field = index_field + 1 @@ -401,13 +401,13 @@ def stencil_ijk( in_field: Field[np.float64], # type: ignore out_field: Field[np.float64], # type: ignore index_field: Field[int], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, 1] + in_field[0, 0, index_field + 1] @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_variable_offsets_and_while_loop(backend): +def test_variable_offsets_and_while_loop(backend) -> None: @gtscript.stencil(backend=backend) def stencil( pe1: Field[np.float64], # type: ignore @@ -415,7 +415,7 @@ def stencil( qin: Field[np.float64], # type: ignore qout: Field[np.float64], # type: ignore lev: Field[gtscript.IJ, np.int_], # type: ignore - ): + ) -> None: with computation(FORWARD), interval(0, -1): if pe2[0, 0, 1] <= pe1[0, 0, lev]: qout = qin[0, 0, 1] @@ -428,7 +428,7 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_nested_while_loop(backend): +def test_nested_while_loop(backend) -> None: @gtscript.stencil(backend=backend) def stencil(field_a: Field[np.float64], field_b: Field[np.int_]): # type: ignore with computation(PARALLEL), interval(...): @@ -440,9 +440,9 @@ def stencil(field_a: Field[np.float64], field_b: Field[np.int_]): # type: ignor @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_mask_with_offset_written_in_conditional(backend): +def test_mask_with_offset_written_in_conditional(backend) -> None: @gtscript.stencil(backend) - def stencil(outp: Field[np.float64]): # type: ignore + def stencil(outp: Field[np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): cond = True if cond[0, -1, 0] or cond[0, 0, 0]: @@ -461,14 +461,14 @@ def stencil(outp: Field[np.float64]): # type: ignore @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_write_data_dim_indirect_addressing(backend): +def test_write_data_dim_indirect_addressing(backend) -> None: INT32_VEC2 = (np.int32, (2,)) def stencil( input_field: Field[gtscript.IJK, np.int32], # type: ignore output_field: Field[gtscript.IJK, INT32_VEC2], # type: ignore index: int, - ): + ) -> None: with computation(PARALLEL), interval(...): output_field[0, 0, 0][index] = input_field @@ -486,14 +486,14 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_read_data_dim_indirect_addressing(backend): +def test_read_data_dim_indirect_addressing(backend) -> None: INT32_VEC2 = (np.int32, (2,)) def stencil( input_field: Field[gtscript.IJK, INT32_VEC2], # type: ignore output_field: Field[gtscript.IJK, np.int32], # type: ignore index: int, - ): + ) -> None: with computation(PARALLEL), interval(...): output_field[0, 0, 0] = input_field[0, 0, 0][index] @@ -512,12 +512,12 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) class TestNegativeOrigin: - def test_negative_origin_i(self, backend): + def test_negative_origin_i(self, backend) -> None: @gtscript.stencil(backend=backend) def stencil_i( input_field: Field[gtscript.IJK, np.int32], # type: ignore output_field: Field[gtscript.IJK, np.int32], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): output_field[0, 0, 0] = input_field[1, 0, 0] @@ -531,12 +531,12 @@ def stencil_i( stencil_i(input_field, output_field, origin={"input_field": (-1, 0, 0)}) assert output_field[0, 0, 0] == 1 - def test_negative_origin_k(self, backend): + def test_negative_origin_k(self, backend) -> None: @gtscript.stencil(backend=backend) def stencil_k( input_field: Field[gtscript.IJK, np.int32], # type: ignore output_field: Field[gtscript.IJK, np.int32], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): output_field[0, 0, 0] = input_field[0, 0, 1] @@ -552,9 +552,9 @@ def stencil_k( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_origin_k_fields(backend): +def test_origin_k_fields(backend) -> None: @gtscript.stencil(backend=backend, rebuild=True) - def k_to_ijk(outp: Field[np.float64], inp: Field[gtscript.K, np.float64]): # type: ignore + def k_to_ijk(outp: Field[np.float64], inp: Field[gtscript.K, np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): outp[0, 0, 0] = inp @@ -580,7 +580,7 @@ def k_to_ijk(outp: Field[np.float64], inp: Field[gtscript.K, np.float64]): # ty @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_tmp_stencil(backend): +def test_tmp_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(6, 6, 6), aligned_index=(0, 0, 0) ) @@ -589,7 +589,10 @@ def test_tmp_stencil(backend): ) @gtscript.stencil(backend=backend) - def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.float64]): # type: ignore + def stencil( + field_in: gtscript.Field[np.float64], # type: ignore + field_out: gtscript.Field[np.float64], # type: ignore + ) -> None: with computation(PARALLEL): with interval(...): tmp = field_in + 1 @@ -610,7 +613,7 @@ def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.f @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_backward_stencil(backend): +def test_backward_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -619,7 +622,10 @@ def test_backward_stencil(backend): ) @gtscript.stencil(backend=backend) - def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.float64]): # type: ignore + def stencil( + field_in: gtscript.Field[np.float64], # type: ignore + field_out: gtscript.Field[np.float64], # type: ignore + ) -> None: with computation(BACKWARD): with interval(-1, None): field_in = 2 @@ -638,7 +644,7 @@ def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.f @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_while_stencil(backend): +def test_while_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(6, 6, 6), aligned_index=(0, 0, 0) ) @@ -647,7 +653,10 @@ def test_while_stencil(backend): ) @gtscript.stencil(backend=backend) - def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.float64]): # type: ignore + def stencil( + field_in: gtscript.Field[np.float64], # type: ignore + field_out: gtscript.Field[np.float64], # type: ignore + ) -> None: with computation(PARALLEL): with interval(...): while field_in < 10: @@ -662,7 +671,7 @@ def stencil(field_in: gtscript.Field[np.float64], field_out: gtscript.Field[np.f @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_higher_dim_literal_stencil(backend): +def test_higher_dim_literal_stencil(backend) -> None: FLOAT64_NDDIM = (np.float64, (4,)) field_in = gt_storage.ones( @@ -677,7 +686,7 @@ def test_higher_dim_literal_stencil(backend): def stencil( vec_field: gtscript.Field[FLOAT64_NDDIM], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): out_field[0, 0, 0] = vec_field[0, 0, 0][2] @@ -689,7 +698,7 @@ def stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_higher_dim_scalar_stencil(backend): +def test_higher_dim_scalar_stencil(backend) -> None: FLOAT64_NDDIM = (np.float64, (4,)) field_in = gt_storage.ones( @@ -705,7 +714,7 @@ def stencil( vec_field: gtscript.Field[FLOAT64_NDDIM], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore scalar_argument: int, - ): + ) -> None: with computation(PARALLEL), interval(...): out_field[0, 0, 0] = vec_field[0, 0, 0][scalar_argument] @@ -735,7 +744,7 @@ def data_dims_with_numpy_int_type( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_native_function_call_stencil(backend): +def test_native_function_call_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -747,7 +756,7 @@ def test_native_function_call_stencil(backend): def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, 0] + sin(0.848062) @@ -757,7 +766,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_unary_operator_stencil(backend): +def test_unary_operator_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -769,7 +778,7 @@ def test_unary_operator_stencil(backend): def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): out_field[0, 0, 0] = -in_field[0, 0, 0] @@ -779,7 +788,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_ternary_operator_stencil(backend): +def test_ternary_operator_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -792,7 +801,7 @@ def test_ternary_operator_stencil(backend): def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, 0] if in_field > 10 else in_field[0, 0, 0] + 1 @@ -804,7 +813,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_mask_stencil(backend): +def test_mask_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -817,7 +826,7 @@ def test_mask_stencil(backend): def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): if in_field[0, 0, 0] > 0: out_field[0, 0, 0] = in_field @@ -831,7 +840,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_k_offset_stencil(backend): +def test_k_offset_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -846,7 +855,7 @@ def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore scalar_value: int, - ): + ) -> None: with computation(PARALLEL), interval(1, None): out_field[0, 0, 0] = in_field[0, 0, scalar_value] @@ -857,7 +866,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_k_offset_field_stencil(backend): +def test_k_offset_field_stencil(backend) -> None: field_in = gt_storage.ones( dtype=np.float64, backend=backend, shape=(4, 4, 4), aligned_index=(0, 0, 0) ) @@ -873,7 +882,7 @@ def test_stencil( in_field: gtscript.Field[np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore idx_field: gtscript.Field[gtscript.IJ, np.int64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(1, None): out_field[0, 0, 0] = in_field[0, 0, idx_field + 1] @@ -884,7 +893,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_k_only_access_stencil(backend): +def test_k_only_access_stencil(backend) -> None: field_in = gt_storage.from_array( np.array([2, 3, 4, 5]), dtype=np.float64, backend=backend, aligned_index=(0,) ) @@ -896,7 +905,7 @@ def test_k_only_access_stencil(backend): def test_stencil( in_field: gtscript.Field[gtscript.K, np.float64], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL): with interval(0, 1): out_field[0, 0, 0] = in_field[1] @@ -910,7 +919,7 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_table_access_stencil(backend): +def test_table_access_stencil(backend) -> None: table_view = gt_storage.from_array( np.array([2, 3, 4, 5]), dtype=np.float64, backend=backend, aligned_index=(0,) ) @@ -922,7 +931,7 @@ def test_table_access_stencil(backend): def test_stencil( table_view: gtscript.GlobalTable[(np.float64, (4))], # type: ignore out_field: gtscript.Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL): with interval(0, 1): out_field[0, 0, 0] = table_view.A[1] @@ -936,9 +945,9 @@ def test_stencil( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_pruned_args_match(backend): +def test_pruned_args_match(backend) -> None: @gtscript.stencil(backend=backend) - def test(out: Field[np.float64], inp: Field[np.float64]): # type: ignore + def test(out: Field[np.float64], inp: Field[np.float64]) -> None: # type: ignore with computation(PARALLEL), interval(...): out = 0.0 with horizontal(region[I[0] - 1, J[0] - 1]): @@ -963,7 +972,7 @@ def test_K_offset_write_simple(backend: str) -> None: # A is untouched # B is written in K+1 and should have K_values, except for the first element (FORWARD) @gtscript.stencil(backend=backend) - def simple(A: Field[np.float64], B: Field[np.float64]): # type: ignore + def simple(A: Field[np.float64], B: Field[np.float64]) -> None: # type: ignore with computation(FORWARD), interval(...): B[0, 0, 1] = A @@ -991,7 +1000,7 @@ def test_K_offset_write_forward(backend: str) -> None: # means while A is update B will have non-updated values of A # Because of the interval, value of B[0] is 0 @gtscript.stencil(backend=backend) - def forward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64): # type: ignore + def forward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64) -> None: # type: ignore with computation(FORWARD), interval(1, None): A[0, 0, -1] = scalar B[0, 0, 0] = A @@ -1022,7 +1031,7 @@ def test_K_offset_write_backward(backend: str) -> None: # means A is update B will get the updated values of A # Because of the interval, B[0] is never written @gtscript.stencil(backend=backend) - def backward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64): # type: ignore + def backward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64) -> None: # type: ignore with computation(BACKWARD), interval(-1, None): A = scalar @@ -1046,13 +1055,17 @@ def backward(A: Field[np.float64], B: Field[np.float64], scalar: np.float64): # @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_K_offset_write_conditional(backend): +def test_K_offset_write_conditional(backend) -> None: arraylib = get_array_library(backend) array_shape = (1, 1, 4) K_values = arraylib.arange(start=40, stop=44) @gtscript.stencil(backend=backend) - def column_physics_conditional(A: Field[np.float64], B: Field[np.float64], scalar: np.float64): # type: ignore + def column_physics_conditional( + A: Field[np.float64], # type: ignore + B: Field[np.float64], # type: ignore + scalar: np.float64, + ) -> None: with computation(BACKWARD), interval(1, -1): if A > 0 and B > 0: A[0, 0, -1] = scalar @@ -1112,11 +1125,11 @@ def column_physics_conditional(A: Field[np.float64], B: Field[np.float64], scala @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_direct_datadims_index(backend): +def test_direct_datadims_index(backend) -> None: F64_VEC4 = (np.float64, (2, 2, 2, 2)) @gtscript.stencil(backend=backend) - def test(out: Field[np.float64], inp: GlobalTable[F64_VEC4]): # type: ignore + def test(out: Field[np.float64], inp: GlobalTable[F64_VEC4]) -> None: # type: ignore with computation(PARALLEL), interval(...): out[0, 0, 0] = inp.A[1, 0, 1, 0] @@ -1128,7 +1141,7 @@ def test(out: Field[np.float64], inp: GlobalTable[F64_VEC4]): # type: ignore @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_function_inline_in_while(backend): +def test_function_inline_in_while(backend) -> None: @gtscript.function def add_42(v): return v + 42 @@ -1137,7 +1150,7 @@ def add_42(v): def test( in_field: Field[np.float64], # type: ignore out_field: Field[np.float64], # type: ignore - ): + ) -> None: with computation(PARALLEL), interval(...): count = 1 while count < 10: @@ -1153,21 +1166,21 @@ def test( @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_cast_in_index(backend): +def test_cast_in_index(backend) -> None: @gtscript.stencil(backend) def cast_in_index( in_field: Field[np.float64], # type: ignore i32: np.int32, i64: np.int64, out_field: Field[np.float64], # type: ignore - ): + ) -> None: """Simple copy stencil with forced cast in index calculation.""" with computation(PARALLEL), interval(...): out_field[0, 0, 0] = in_field[0, 0, i32 - i64] @pytest.mark.parametrize("backend", ALL_BACKENDS) -def test_read_after_write_stencil(backend): +def test_read_after_write_stencil(backend) -> None: """Stencil with multiple read after write access patterns.""" @gtscript.stencil(backend=backend) @@ -1181,7 +1194,7 @@ def lagrangian_contributions( q4_4: Field[np.float64], # type: ignore dp1: Field[np.float64], # type: ignore lev: Field[gtscript.IJ, np.int64], # type: ignore - ): + ) -> None: """ Args: q (out): @@ -1241,9 +1254,9 @@ def lagrangian_contributions( for backend in ["gt:cpu_ifirst", "numpy"] ], ) -def test_absolute_K_index_raise(backend): +def test_absolute_K_index_raise(backend) -> None: @gtscript.stencil(backend=backend) - def test_absolute_k_index(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: + def test_absolute_k_index(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore with computation(PARALLEL), interval(...): out_field = in_field.at(K=2) @@ -1256,7 +1269,7 @@ def test_absolute_k_index(in_field: Field[np.float64], out_field: Field[np.float pytest.param("dace:gpu", marks=[pytest.mark.uses_dace, pytest.mark.requires_gpu]), ], ) -def test_absolute_K_index(backend): +def test_absolute_K_index(backend) -> None: domain = (5, 5, 5) in_arr = gt_storage.ones(backend=backend, shape=domain, dtype=np.float64) @@ -1266,7 +1279,7 @@ def test_absolute_K_index(backend): out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=np.float64) @gtscript.stencil(backend=backend) - def test_literal_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: + def test_literal_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore with computation(PARALLEL), interval(...): out_field = in_field.at(K=2) @@ -1278,7 +1291,9 @@ def test_literal_access(in_field: Field[np.float64], out_field: Field[np.float64 @gtscript.stencil(backend=backend) def test_parameter_access( - in_field: Field[np.float64], out_field: Field[np.float64], idx: int + in_field: Field[np.float64], # type:ignore + out_field: Field[np.float64], # type:ignore + idx: int, ) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=idx) @@ -1290,9 +1305,9 @@ def test_parameter_access( assert (out_arr[:, :, :] == 42.42).all() @gtscript.stencil(backend=backend, externals={"K4": 4}) - def test_external_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: + def test_external_access(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore with computation(PARALLEL), interval(...): - from __externals__ import K4 + from __externals__ import K4 # type:ignore out_field = in_field.at(K=K4) @@ -1304,7 +1319,9 @@ def test_external_access(in_field: Field[np.float64], out_field: Field[np.float6 @gtscript.stencil(backend=backend) def test_field_access( - in_field: Field[np.float64], index_field: Field[IJ, np.int64], out_field: Field[np.float64] + in_field: Field[np.float64], # type:ignore + index_field: Field[IJ, np.int64], # type:ignore + out_field: Field[np.float64], # type:ignore ) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=index_field) @@ -1318,7 +1335,9 @@ def test_field_access( @gtscript.stencil(backend=backend) def test_field_access_computation( - in_field: Field[np.float64], index_field: Field[IJ, np.int32], out_field: Field[np.float64] + in_field: Field[np.float64], # type:ignore + index_field: Field[IJ, np.int32], # type:ignore + out_field: Field[np.float64], # type:ignore ) -> None: with computation(PARALLEL), interval(...): out_field = in_field.at(K=index_field - 1) @@ -1332,8 +1351,8 @@ def test_field_access_computation( @gtscript.stencil(backend=backend) def test_lower_dim_field( - k_field: Field[K, np.float64], - out_field: Field[np.float64], + k_field: Field[K, np.float64], # type:ignore + out_field: Field[np.float64], # type:ignore ) -> None: with computation(PARALLEL), interval(...): out_field = k_field.at(K=2) @@ -1346,8 +1365,8 @@ def test_lower_dim_field( @gtscript.stencil(backend=backend) def test_conditional_absolute( - in_field: Field[np.float64], - out_field: Field[np.float64], + in_field: Field[np.float64], # type:ignore + out_field: Field[np.float64], # type:ignore ) -> None: with computation(PARALLEL), interval(...): k_level = 0 @@ -1375,7 +1394,9 @@ def test_iterator_access(backend: str) -> None: @gtscript.stencil(backend=backend) def test_all_valid_usage( - field_A: Field[np.float64], field_B: Field[np.float64], offsets: Field[K, np.int32] + field_A: Field[np.float64], # type:ignore + field_B: Field[np.float64], # type:ignore + offsets: Field[K, np.int32], # type:ignore ) -> None: with computation(PARALLEL), interval(...): if K == 2: @@ -1408,7 +1429,7 @@ def test_iterator_access_raises_in_unsupported_backends(backend: str) -> None: field_B = gt_storage.zeros(backend=backend, shape=domain, dtype=np.float64) @gtscript.stencil(backend=backend) - def test_all_valid_usage(field_A: Field[np.float64], field_B: Field[np.float64]) -> None: + def test_all_valid_usage(field_A: Field[np.float64], field_B: Field[np.float64]) -> None: # type:ignore with computation(PARALLEL), interval(...): if K == 2: field_A = 20.20 @@ -1417,7 +1438,7 @@ def test_all_valid_usage(field_A: Field[np.float64], field_B: Field[np.float64]) test_all_valid_usage(field_A, field_B) -def test_runtime_interval_bounds(): +def test_runtime_interval_bounds() -> None: backend = "debug" domain = (5, 5, 10) in_arr = gt_storage.ones(backend=backend, shape=domain, dtype=np.float64) @@ -1522,14 +1543,14 @@ def test_temporary( ), ], ) -def test_runtime_interval_raises(backend): +def test_runtime_interval_raises(backend) -> None: @gtscript.stencil(backend=backend) def test_stencil( out_field: Field[np.float64], # type: ignore input_data: Field[np.float64], # type: ignore index_data: Field[gtscript.IJ, np.int64], # type: ignore scalar_arg: int, - ): + ) -> None: with computation(FORWARD), interval(0, 1): temporary: Field[IJ, np.float64] = 7 # type: ignore @@ -1552,16 +1573,16 @@ def test_stencil( pytest.param("dace:gpu", marks=[pytest.mark.uses_dace, pytest.mark.requires_gpu]), ], ) -def test_2d_temporaries(backend): +def test_2d_temporaries(backend) -> None: domain = (5, 5, 3) in_arr = gt_storage.ones(backend=backend, shape=domain, dtype=np.float64) out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=np.float64) @gtscript.stencil(backend=backend) - def test_with_plain_gt4py(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: + def test_with_plain_gt4py(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore with computation(FORWARD), interval(0, 1): - tmp_2D: Field[IJ, np.float64] = 0 + tmp_2D: Field[IJ, np.float64] = 0 # type:ignore with computation(FORWARD), interval(...): tmp_2D = tmp_2D + in_field @@ -1575,9 +1596,9 @@ def test_with_plain_gt4py(in_field: Field[np.float64], out_field: Field[np.float assert (out_arr[:, :, :] == domain[2]).all() @gtscript.stencil(backend=backend, dtypes={"MyFancySymbol": Field[IJ, np.float64]}) - def test_with_user_dtype(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: + def test_with_user_dtype(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore with computation(FORWARD), interval(0, 1): - tmp_2D: MyFancySymbol = 0 + tmp_2D: MyFancySymbol = 0 # type:ignore with computation(FORWARD), interval(...): out_field = tmp_2D @@ -1588,10 +1609,11 @@ def test_with_user_dtype(in_field: Field[np.float64], out_field: Field[np.float6 @gtscript.stencil(backend=backend) def test_failing_on_non_IJ( - in_field: Field[np.float64], out_field: Field[np.float64] + in_field: Field[np.float64], # type:ignore + out_field: Field[np.float64], # type:ignore ) -> None: with computation(FORWARD), interval(0, 1): - tmp_2D: Field[K, np.float64] = 0 + tmp_2D: Field[K, np.float64] = 0 # type:ignore with computation(FORWARD), interval(...): out_field = tmp_2D @@ -1609,11 +1631,11 @@ def test_failing_on_non_IJ( for backend in ["gt:cpu_ifirst", "gt:cpu_kfirst", "gt:gpu"] ], ) -def test_2d_temporaries_raises(backend): +def test_2d_temporaries_raises(backend) -> None: @gtscript.stencil(backend=backend) - def test_with_user_dtypes(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: + def test_with_user_dtypes(in_field: Field[np.float64], out_field: Field[np.float64]) -> None: # type:ignore with computation(FORWARD), interval(0, 1): - tmp_2D: Field[IJ, np.float64] = 0 + tmp_2D: Field[IJ, np.float64] = 0 # type:ignore with computation(FORWARD), interval(...): out_field = tmp_2D @@ -1629,7 +1651,9 @@ def test_upcasting_both_sides_of_assignment(backend: str) -> None: @gtscript.stencil(backend=backend) def test_upcasting_stencil( - in_field: Field[np.float64], index_field: Field[IJ, np.int32], out_field: Field[np.float64] + in_field: Field[np.float64], # type:ignore + index_field: Field[IJ, np.int32], # type:ignore + out_field: Field[np.float64], # type:ignore ) -> None: with computation(FORWARD), interval(...): out_field[0, 0, index_field - 1] = in_field @@ -1638,7 +1662,7 @@ def test_upcasting_stencil( assert (input == output).all() -@pytest.mark.parametrize("backend", ("debug",)) # ALL_BACKENDS) +@pytest.mark.parametrize("backend", ALL_BACKENDS) def test_upcasting_leave_integer_power_arguments_alone(backend: str) -> None: domain = (5, 5, 5) @@ -1650,7 +1674,9 @@ def test_upcasting_leave_integer_power_arguments_alone(backend: str) -> None: @gtscript.stencil(backend=backend) def test_upcasting_stencil( - in_field: Field[np.float32], squared: Field[IJ, np.int32], out_field: Field[np.float32] + in_field: Field[np.float32], # type:ignore + squared: Field[IJ, np.int32], # type:ignore + out_field: Field[np.float32], # type:ignore ) -> None: with computation(FORWARD), interval(...): out_field = in_field**squared @@ -1662,14 +1688,14 @@ def test_no_write_and_read_with_horizontal_offset() -> None: with pytest.raises(ValueError, match="Self-assignment with offset in I or J is illegal."): @gtscript.stencil(backend="debug") - def self_assign_offset(field: Field[np.float64]) -> None: + def self_assign_offset(field: Field[np.float64]) -> None: # type:ignore with computation(PARALLEL), interval(...): field = (field[I - 1] + field[I + 1]) / 2 with pytest.raises(ValueError, match="Illegal write and read with horizontal offset"): @gtscript.stencil(backend="debug") - def self_assign_offset(field: Field[np.float64]) -> None: + def self_assign_offset(field: Field[np.float64]) -> None: # type:ignore with computation(PARALLEL), interval(...): tmp = (field[J - 1] + field[J + 1]) / 2 field = tmp * 2 @@ -1679,14 +1705,14 @@ def test_k_offsets_in_parallel_loops() -> None: with pytest.raises(ValueError, match="write and read with k-offsets in PARALLEL"): @gtscript.stencil(backend="debug") - def self_assign_offset_parallel(field: Field[np.int32]) -> None: + def self_assign_offset_parallel(field: Field[np.int32]) -> None: # type:ignore with computation(PARALLEL), interval(1, None): field = field[K - 1] * 2 with pytest.raises(ValueError, match="write and read with k-offsets in PARALLEL"): @gtscript.stencil(backend="debug") - def self_assign_offset_parallel_temp(field: Field[np.int32]) -> None: + def self_assign_offset_parallel_temp(field: Field[np.int32]) -> None: # type:ignore with computation(PARALLEL), interval(1, None): tmp = field[K - 1] field = tmp * 2 @@ -1696,7 +1722,7 @@ def self_assign_offset_parallel_temp(field: Field[np.int32]) -> None: ): @gtscript.stencil(backend="debug") - def mixed_read_write(field: Field[np.int32]): + def mixed_read_write(field: Field[np.int32]) -> None: # type:ignore with computation(PARALLEL), interval(...): level = field.at(K=1) field = 2 * level @@ -1706,31 +1732,31 @@ def mixed_read_write(field: Field[np.int32]): ): @gtscript.stencil(backend="debug") - def mixed_read_write(field: Field[np.int32], offset: int = -1): + def mixed_read_write(field: Field[np.int32], offset: int = -1) -> None: # type:ignore with computation(PARALLEL), interval(1, None): bottom = field[0, 0, offset] field = field + 2 * bottom # center reads and writes are allowed @gtscript.stencil(backend="debug") - def self_assignment_center_read_parallel(field: Field[np.int32]) -> None: + def self_assignment_center_read_parallel(field: Field[np.int32]) -> None: # type:ignore with computation(PARALLEL), interval(...): field = field[0, 0, 0] * 2 @gtscript.stencil(backend="debug") - def self_assignment_center_write_parallel(field: Field[np.int32]) -> None: + def self_assignment_center_write_parallel(field: Field[np.int32]) -> None: # type:ignore with computation(PARALLEL), interval(...): field[0, 0, 0] = field * 2 # not mixing reads and writes are allowed (e.g. index fields) @gtscript.stencil(backend="debug") - def self_assignment_center_parallel(field: Field[np.float32], index: Field[np.int32]) -> None: + def self_assignment_center_parallel(field: Field[np.float32], index: Field[np.int32]) -> None: # type:ignore with computation(PARALLEL), interval(1, None): field = index + index[K - 1] * 2 # parallel intervals of static size 1 are allowed @gtscript.stencil(backend="debug") - def the_stencil(field: Field[np.bool_]) -> None: + def the_stencil(field: Field[np.bool_]) -> None: # type:ignore with computation(PARALLEL): with interval(0, 1): field = field[K + 1] @@ -1741,12 +1767,12 @@ def the_stencil(field: Field[np.bool_]) -> None: @pytest.mark.parametrize("backend", ALL_BACKENDS) def test_self_assignment_in_forward(backend: str) -> None: @gtscript.stencil(backend=backend) - def self_assignment_parallel(field: Field[np.int32]) -> None: + def self_assignment_parallel(field: Field[np.int32]) -> None: # type:ignore with computation(FORWARD), interval(1, None): field = field[K - 1] * 2 @gtscript.stencil(backend=backend) - def self_assignment_2_parallel(field: Field[np.int32]) -> None: + def self_assignment_2_parallel(field: Field[np.int32]) -> None: # type:ignore with computation(FORWARD), interval(1, None): tmp = field[K - 1] field = tmp * 2 @@ -1762,7 +1788,9 @@ def test_reset_mask_2d(backend: str) -> None: @gtscript.stencil(backend=backend) def test_set_2d_mask( - dp1: Field[np.float64], pe1: Field[np.float64], lev: Field[IJ, np.int32] + dp1: Field[np.float64], # type:ignore + pe1: Field[np.float64], # type:ignore + lev: Field[IJ, np.int32], # type:ignore ) -> None: with computation(PARALLEL), interval(0, -1): dp1 = pe1[0, 0, 1] - pe1 From 00e99994cfefb0f65efb32395138cf6327d63b15 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 23 Sep 2026 10:37:52 +0200 Subject: [PATCH 42/48] Reapply "feat[cartesian]: DaCe optimal for/map schedule (#2628)" This reverts commit e5805fd61b0f7827112eb6bcdae39f4d6ad84682. # Conflicts: # src/gt4py/cartesian/gtc/dace/oir_to_treeir.py --- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 48 ++++---- .../stencil_definitions.py | 2 +- .../test_code_generation.py | 107 ++++++++++++++++++ 3 files changed, 129 insertions(+), 28 deletions(-) diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index 8d94b5e093..4c0b2c78c8 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -31,10 +31,8 @@ """Default dace residency types per device type.""" -def _resolve_default_map_schedule( - device_type: dtypes.DeviceType, -) -> dtypes.ScheduleType: - """Default kernel target per device type.""" +def _resolve_map_schedule(device_type: dtypes.DeviceType) -> dtypes.ScheduleType: + """Optimal kernel schedule type based on target device.""" if device_type == dtypes.DeviceType.GPU: return dtypes.ScheduleType.GPU_Device @@ -44,7 +42,7 @@ def _resolve_default_map_schedule( if not gt_config.build_settings["openmp"]["use_openmp"]: return dtypes.ScheduleType.Sequential - return dtypes.ScheduleType.Default + return dtypes.ScheduleType.CPU_Multicore class OIRToTreeIR(eve.NodeVisitor): @@ -148,7 +146,7 @@ def visit_HorizontalExecution(self, node: oir.HorizontalExecution, ctx: tir.Cont loop = tir.HorizontalLoop( bounds_i=tir.Bounds(start=axis_start_i, end=axis_end_i), bounds_j=tir.Bounds(start=axis_start_j, end=axis_end_j), - schedule=_resolve_default_map_schedule(self._device_type), + schedule=_resolve_map_schedule(self._device_type), children=[], parent=ctx.current_scope, ) @@ -283,19 +281,6 @@ def visit_Interval( return tir.Bounds(start=start, end=end) - def _vertical_loop_schedule(self) -> dtypes.ScheduleType: - """ - Defines the vertical loop schedule. - - Current strategy is to - - keep the vertical loop on the host for both, CPU and GPU targets - - and run it in parallel on CPU and sequential on GPU. - """ - if self._device_type == dtypes.DeviceType.GPU: - return dtypes.ScheduleType.Sequential - - return _resolve_default_map_schedule(self._device_type) - def visit_VerticalLoopSection( self, node: oir.VerticalLoopSection, ctx: tir.Context, loop_order: common.LoopOrder ) -> None: @@ -306,14 +291,23 @@ def visit_VerticalLoopSection( axis_end=tir.Axis.K.domain_dace_symbol(), ) - loop = tir.VerticalLoop( - iteration_variable=tir.Axis.K.iteration_symbol(), - loop_order=loop_order, - bounds_k=bounds, - schedule=self._vertical_loop_schedule(), - children=[], - parent=ctx.current_scope, - ) + loop: tir.SequentialVerticalLoop | tir.ParallelVerticalLoop + if loop_order == common.LoopOrder.PARALLEL: + loop = tir.ParallelVerticalLoop( + iteration_variable=tir.Axis.K.iteration_symbol(), + bounds_k=bounds, + schedule=_resolve_map_schedule(self._device_type), + children=[], + parent=ctx.current_scope, + ) + else: + loop = tir.SequentialVerticalLoop( + iteration_variable=tir.Axis.K.iteration_symbol(), + bounds_k=bounds, + loop_order=loop_order, + children=[], + parent=ctx.current_scope, + ) with loop.scope(ctx): self.visit(node.horizontal_executions, ctx=ctx) diff --git a/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py b/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py index 3c97675fd5..b3d412cb79 100644 --- a/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py +++ b/tests/cartesian_tests/integration_tests/multi_feature_tests/stencil_definitions.py @@ -78,7 +78,7 @@ def copy_stencil(field_a: Field3D, field_b: Field3D): @gtscript.function def a_gtscript_function(b): - return sqrt(abs(b[0, 1, 0])) + return sqrt(abs(b[0, 0, 0])) @register diff --git a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py index 5519ff790c..f16c92c18a 100644 --- a/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py +++ b/tests/cartesian_tests/integration_tests/multi_feature_tests/test_code_generation.py @@ -24,11 +24,17 @@ J, K, IJ, + IJK, computation, horizontal, interval, region, sin, + tan, + isfinite, + isinf, + isnan, + sqrt, ) from gt4py.storage.cartesian import utils as storage_utils @@ -1800,3 +1806,104 @@ def test_set_2d_mask( test_set_2d_mask(output, input, mask_2d) assert (mask_2d == 0).all() + + +@pytest.mark.parametrize( + "backend", + [ + "debug", + pytest.param("dace:cpu", marks=[pytest.mark.uses_dace]), + pytest.param( + "dace:gpu", + marks=[ + pytest.mark.uses_dace, + pytest.mark.requires_gpu, + pytest.mark.xfail( + raises=SystemExit, + reason="DaCe issue: Missing `_gbar` symbol for global sync inside nested SDFG.", + ), + ], + ), + pytest.param("gt:gpu", marks=[pytest.mark.requires_gpu]), + ], +) +def test_offset_j_in_temporaries(backend: str) -> None: + @gtscript.function + def a_gtscript_function(b): + return sqrt(abs(b[0, 1, 0])) + + @gtscript.stencil(backend=backend) + def test_stencil_offset_j_in_temporaries( + field_in: Field[IJK, np.float64], # type: ignore + field_out: Field[IJK, np.float64], # type: ignore + ) -> None: + with computation(PARALLEL), interval(...): + abs_res = abs(field_in) + tan_res = tan(abs_res) + + # This is an offset in J on a temporary, it will + # require a global kernel sync in KJI for GPU + sqrt_res = a_gtscript_function(tan_res) + + field_out = ( + sqrt_res + if isfinite(sqrt_res) + else field_in + if isinf(sqrt_res) + else field_out + if isnan(sqrt_res) + else 0.0 + ) + + +@gtscript.enum +class MyEnum(IntEnum): + Zero = 0 + A = 10 + B = 20 + C = 30 + + +@pytest.mark.parametrize("backend", ALL_BACKENDS) +def test_enum_runtime(backend): + + @gtscript.stencil(backend=backend) + def the_stencil(out_field: Field[int], order: MyEnum): # type: ignore + with computation(PARALLEL), interval(0, 1): + out_field = 32 + if order < MyEnum.A: + out_field = MyEnum.A + + with computation(PARALLEL), interval(1, 2): + out_field = 23 + out_field = MyEnum.B + + with computation(PARALLEL), interval(2, None): + out_field = 56 + out_field = MyEnum.C + + domain = (5, 5, 5) + out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=int) + + the_stencil(out_arr, MyEnum.Zero) + + assert out_arr[0, 0, 0] == MyEnum.A.value + assert out_arr[0, 0, 1] == MyEnum.B.value + assert (out_arr[0, 0, 2:] == MyEnum.C.value).all() + + +@pytest.mark.parametrize("backend", ALL_BACKENDS) +def test_negated_bool_runtime(backend): + + @gtscript.stencil(backend=backend) + def the_stencil(out_field: Field[int], done: bool): # type: ignore + with computation(PARALLEL), interval(...): + if not done: + out_field = 1 + + domain = (5, 5, 5) + out_arr = gt_storage.zeros(backend=backend, shape=domain, dtype=int) + + the_stencil(out_arr, done=False) + + assert (out_arr[:] == 1).all() From 5ad3321cb2580d7c7f843cd4faef9fc8bd21f9bd Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 23 Sep 2026 11:01:16 +0200 Subject: [PATCH 43/48] Parametrize parallel vertical loop schedule --- src/gt4py/cartesian/caching.py | 3 +++ src/gt4py/cartesian/config.py | 6 +++++- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 10 +++++++++- 3 files changed, 17 insertions(+), 2 deletions(-) diff --git a/src/gt4py/cartesian/caching.py b/src/gt4py/cartesian/caching.py index 62f950dc39..00a119aaff 100644 --- a/src/gt4py/cartesian/caching.py +++ b/src/gt4py/cartesian/caching.py @@ -323,6 +323,9 @@ def stencil_id(self) -> StencilID: fingerprint["dace_version"] = dace_version if self.builder.backend.name == "dace:gpu": fingerprint["default_block_size"] = gt_config.DACE_DEFAULT_BLOCK_SIZE + fingerprint["parallel_vertical_loop_schedule"] = ( + gt_config.DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE + ) # ignore type because attrclass StencilID has generated constructor return StencilID( # type: ignore diff --git a/src/gt4py/cartesian/config.py b/src/gt4py/cartesian/config.py index 03525caf21..ed8a82b788 100644 --- a/src/gt4py/cartesian/config.py +++ b/src/gt4py/cartesian/config.py @@ -8,7 +8,7 @@ import multiprocessing import os -from typing import Any +from typing import Any, Literal import gridtools_cpp @@ -92,3 +92,7 @@ os.environ.setdefault("DACE_CONFIG", os.path.join(os.path.abspath("."), ".dace.conf")) DACE_DEFAULT_BLOCK_SIZE = os.environ.get("DACE_DEFAULT_BLOCK_SIZE", "64,8,1") + +DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE: Literal[ + "sequential", "gpu_device", "gpu_threadblock", "gpu_threadblock_dynamic" +] = os.environ.get("DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE", "gpu_device") # type: ignore[assignment] diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index 4c0b2c78c8..b0be3606db 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -34,7 +34,15 @@ def _resolve_map_schedule(device_type: dtypes.DeviceType) -> dtypes.ScheduleType: """Optimal kernel schedule type based on target device.""" if device_type == dtypes.DeviceType.GPU: - return dtypes.ScheduleType.GPU_Device + match gt_config.DACE_DEFAULT_BLOCK_SIZE: + case "gpu_device": + return dtypes.ScheduleType.GPU_Device + case "gpu_threadblock": + return dtypes.ScheduleType.GPU_ThreadBlock + case "gpu_threadblock_dynamic": + return dtypes.ScheduleType.GPU_ThreadBlock_Dynamic + case _: + return dtypes.ScheduleType.Sequential if device_type != dtypes.DeviceType.CPU: raise NotImplementedError(f"Schedule Tree bridge does not support {device_type}") From e031aa7bb319f3ca1806820cf88a6b3b9957a86c Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 23 Sep 2026 11:40:22 +0200 Subject: [PATCH 44/48] Apply map collapse transformation to sdgf --- src/gt4py/cartesian/backend/dace_backend.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/gt4py/cartesian/backend/dace_backend.py b/src/gt4py/cartesian/backend/dace_backend.py index bd4b456349..9452e4f5aa 100644 --- a/src/gt4py/cartesian/backend/dace_backend.py +++ b/src/gt4py/cartesian/backend/dace_backend.py @@ -18,6 +18,7 @@ from dace.codegen import codeobject from dace.sdfg.analysis.schedule_tree import treenodes as tn from dace.sdfg.utils import inline_sdfgs +from dace.transformation.dataflow import MapCollapse from gt4py._core import definitions as core_defs from gt4py.cartesian import config as gt_config, definitions @@ -440,6 +441,7 @@ def sdfg_via_schedule_tree(self, *, validate: bool = False, simplify: bool = Tru # - `LiftTrivialIf` because it's dead slow (e.g. fv3 acoustics parsing takes >90min compared to 10-15min without) skip={"ScalarToSymbolPromotion", "ControlFlowRaising", "LiftTrivialIf"}, ) + sdfg.apply_transformations_repeated(MapCollapse, progress=False, validate=validate) if do_cache: self._save_sdfg(sdfg, path) From b41448f15a62fc8898ef6134dbb1eb9ab27fd424 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 23 Sep 2026 14:03:21 +0200 Subject: [PATCH 45/48] Fix typo --- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index b0be3606db..8d81171274 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -34,7 +34,7 @@ def _resolve_map_schedule(device_type: dtypes.DeviceType) -> dtypes.ScheduleType: """Optimal kernel schedule type based on target device.""" if device_type == dtypes.DeviceType.GPU: - match gt_config.DACE_DEFAULT_BLOCK_SIZE: + match gt_config.DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE: case "gpu_device": return dtypes.ScheduleType.GPU_Device case "gpu_threadblock": From 92923b18d0b53637fdef70ef3e5ea4cfcd2dfbff Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 23 Sep 2026 16:12:09 +0200 Subject: [PATCH 46/48] Remove debug code --- src/gt4py/cartesian/caching.py | 12 ++---------- src/gt4py/cartesian/config.py | 6 +----- src/gt4py/cartesian/gtc/dace/oir_to_treeir.py | 10 +--------- 3 files changed, 4 insertions(+), 24 deletions(-) diff --git a/src/gt4py/cartesian/caching.py b/src/gt4py/cartesian/caching.py index 00a119aaff..e3bdb280f3 100644 --- a/src/gt4py/cartesian/caching.py +++ b/src/gt4py/cartesian/caching.py @@ -19,9 +19,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional from cached_property import cached_property -from dace import __version__ as dace_version -from gt4py import __version__ as gt4py_version from gt4py.cartesian import config as gt_config, utils as gt_utils from gt4py.cartesian.definitions import StencilID @@ -301,7 +299,6 @@ def _extract_api_annotations(self) -> List[str]: @property def stencil_id(self) -> StencilID: fingerprint = { - "gt4py_version": gt4py_version, "__main__": self.builder.definition._gtscript_["canonical_ast"], "docstring": inspect.getdoc(self.builder.definition), "api_annotations": f"[{', '.join(self._extract_api_annotations())}]", @@ -319,13 +316,8 @@ def stencil_id(self) -> StencilID: fingerprint["extra_compile_args"] = self.builder.options.backend_opts.get( "extra_compile_args", gt_config.GT4PY_EXTRA_COMPILE_ARGS ) - if "dace" in self.builder.backend.name: - fingerprint["dace_version"] = dace_version - if self.builder.backend.name == "dace:gpu": - fingerprint["default_block_size"] = gt_config.DACE_DEFAULT_BLOCK_SIZE - fingerprint["parallel_vertical_loop_schedule"] = ( - gt_config.DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE - ) + if self.builder.backend.name == "dace:gpu": + fingerprint["default_block_size"] = gt_config.DACE_DEFAULT_BLOCK_SIZE # ignore type because attrclass StencilID has generated constructor return StencilID( # type: ignore diff --git a/src/gt4py/cartesian/config.py b/src/gt4py/cartesian/config.py index ed8a82b788..03525caf21 100644 --- a/src/gt4py/cartesian/config.py +++ b/src/gt4py/cartesian/config.py @@ -8,7 +8,7 @@ import multiprocessing import os -from typing import Any, Literal +from typing import Any import gridtools_cpp @@ -92,7 +92,3 @@ os.environ.setdefault("DACE_CONFIG", os.path.join(os.path.abspath("."), ".dace.conf")) DACE_DEFAULT_BLOCK_SIZE = os.environ.get("DACE_DEFAULT_BLOCK_SIZE", "64,8,1") - -DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE: Literal[ - "sequential", "gpu_device", "gpu_threadblock", "gpu_threadblock_dynamic" -] = os.environ.get("DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE", "gpu_device") # type: ignore[assignment] diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py index 8d81171274..4c0b2c78c8 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_treeir.py @@ -34,15 +34,7 @@ def _resolve_map_schedule(device_type: dtypes.DeviceType) -> dtypes.ScheduleType: """Optimal kernel schedule type based on target device.""" if device_type == dtypes.DeviceType.GPU: - match gt_config.DACE_PARALLEL_VERTICAL_LOOP_SCHEDULE: - case "gpu_device": - return dtypes.ScheduleType.GPU_Device - case "gpu_threadblock": - return dtypes.ScheduleType.GPU_ThreadBlock - case "gpu_threadblock_dynamic": - return dtypes.ScheduleType.GPU_ThreadBlock_Dynamic - case _: - return dtypes.ScheduleType.Sequential + return dtypes.ScheduleType.GPU_Device if device_type != dtypes.DeviceType.CPU: raise NotImplementedError(f"Schedule Tree bridge does not support {device_type}") From b9926f478d7cea07871c9a5a4a51007ded6084c0 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 30 Sep 2026 15:05:22 +0200 Subject: [PATCH 47/48] Do not validate map-collapse pass --- src/gt4py/cartesian/backend/dace_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/cartesian/backend/dace_backend.py b/src/gt4py/cartesian/backend/dace_backend.py index 9452e4f5aa..3bd76bcfeb 100644 --- a/src/gt4py/cartesian/backend/dace_backend.py +++ b/src/gt4py/cartesian/backend/dace_backend.py @@ -441,7 +441,7 @@ def sdfg_via_schedule_tree(self, *, validate: bool = False, simplify: bool = Tru # - `LiftTrivialIf` because it's dead slow (e.g. fv3 acoustics parsing takes >90min compared to 10-15min without) skip={"ScalarToSymbolPromotion", "ControlFlowRaising", "LiftTrivialIf"}, ) - sdfg.apply_transformations_repeated(MapCollapse, progress=False, validate=validate) + sdfg.apply_transformations_repeated(MapCollapse, progress=False, validate=False) if do_cache: self._save_sdfg(sdfg, path) From 19686d1766f181871486d3e5671b835d65deac28 Mon Sep 17 00:00:00 2001 From: stubbiali Date: Wed, 30 Sep 2026 15:06:19 +0200 Subject: [PATCH 48/48] Fix unit tests --- .../backend_tests/test_dace_backend.py | 32 +++++++++++++------ 1 file changed, 22 insertions(+), 10 deletions(-) diff --git a/tests/cartesian_tests/unit_tests/backend_tests/test_dace_backend.py b/tests/cartesian_tests/unit_tests/backend_tests/test_dace_backend.py index 32ba6659cc..96b2546778 100644 --- a/tests/cartesian_tests/unit_tests/backend_tests/test_dace_backend.py +++ b/tests/cartesian_tests/unit_tests/backend_tests/test_dace_backend.py @@ -113,8 +113,11 @@ def test_dace_cpu_loop_structure(): state = sdfg.states()[0] loop_indices = [node.map.params for node in state.nodes() if isinstance(node, nodes.MapEntry)] - assert loop_indices[0] == ["__k"] - assert loop_indices[1] == ["__i", "__j"] + assert loop_indices[0] == [ + Axis.K.iteration_symbol(), + Axis.I.iteration_symbol(), + Axis.J.iteration_symbol(), + ] def test_dace_cpu_kfirst_loop_structure(): @@ -125,8 +128,11 @@ def test_dace_cpu_kfirst_loop_structure(): state = sdfg.states()[0] loop_indices = [node.map.params for node in state.nodes() if isinstance(node, nodes.MapEntry)] - assert loop_indices[0] == ["__i", "__j"] - assert loop_indices[1] == ["__k"] + assert loop_indices[0] == [ + Axis.I.iteration_symbol(), + Axis.J.iteration_symbol(), + Axis.K.iteration_symbol(), + ] builder = StencilBuilder(copy_forward_stencil, backend="dace:cpu_kfirst") manager = SDFGManager(builder) @@ -164,8 +170,11 @@ def test_dace_cpu_KJI_loop_structure(): loop_indices = [ node.map.params for node in state.nodes() if isinstance(node, nodes.MapEntry) ] - assert loop_indices[0] == ["__k"] - assert loop_indices[1] == ["__j", "__i"] + assert loop_indices[0] == [ + Axis.K.iteration_symbol(), + Axis.J.iteration_symbol(), + Axis.I.iteration_symbol(), + ] builder = StencilBuilder(copy_forward_stencil, backend="dace:cpu_KJI") manager = SDFGManager(builder) @@ -174,12 +183,12 @@ def test_dace_cpu_KJI_loop_structure(): # Expect LoopRegion for K outside loop_region: LoopRegion = list(sdfg.all_control_flow_blocks())[0] - assert loop_region.loop_variable == "__k" + assert loop_region.loop_variable == Axis.K.iteration_symbol() # Expect JI Map and in loop_body state (#2) state = loop_region.start_block assert [node.map.params for node in state.nodes() if isinstance(node, nodes.MapEntry)] == [ - ["__j", "__i"] + [Axis.J.iteration_symbol(), Axis.I.iteration_symbol()], ] @@ -199,7 +208,10 @@ def test_dace_cpu_KJI_loop_structure_parallel(): # Expect a Map for IJ outside map_entry_nodes = [node for node in state.nodes() if isinstance(node, nodes.MapEntry)] assert len(map_entry_nodes) == 1, "expect one MapEntry node" - assert map_entry_nodes[0].map.params == ["__j", "__i"] + assert map_entry_nodes[0].map.params == [ + Axis.J.iteration_symbol(), + Axis.I.iteration_symbol(), + ] # Expect LoopRegion for K inside map nsdfg_nodes = [node for node in state.nodes() if isinstance(node, nodes.NestedSDFG)] @@ -208,4 +220,4 @@ def test_dace_cpu_KJI_loop_structure_parallel(): assert len(for_nested_nodes) == 1 loop_region = for_nested_nodes[0] assert isinstance(loop_region, LoopRegion) - assert loop_region.loop_variable == "__k" + assert loop_region.loop_variable == Axis.K.iteration_symbol()