From c247d4110e4f15b4f76327fa11b138795a75325d Mon Sep 17 00:00:00 2001 From: stubbiali Date: Thu, 3 Oct 2024 18:14:39 +0200 Subject: [PATCH 01/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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/33] 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)