From f86b2c8093b20ef4984ce4921cf829fd69e55017 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Fri, 30 Jan 2026 17:44:57 +0100 Subject: [PATCH 01/49] wip --- src/gt4py/next/iterator/transforms/pass_manager.py | 14 ++++++++++++-- .../runners/dace/workflow/translation.py | 7 ++++++- 2 files changed, 18 insertions(+), 3 deletions(-) diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index 1fb30b096d..32646d27c6 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -54,6 +54,8 @@ def apply_common_transforms( unroll_reduce=False, common_subexpression_elimination=True, force_inline_lambda_args=False, + fuse_maps=True, + transform_concat_where_to_as_fieldop=True, #: A dictionary mapping axes names to their length. See :func:`infer_domain.infer_expr` for #: more details. symbolic_domain_sizes: Optional[dict[str, str]] = None, @@ -92,7 +94,8 @@ def apply_common_transforms( ir = prune_empty_concat_where.prune_empty_concat_where(ir) ir = remove_broadcast.RemoveBroadcast.apply(ir) - ir = concat_where.transform_to_as_fieldop(ir) + if transform_concat_where_to_as_fieldop: + ir = concat_where.transform_to_as_fieldop(ir) for _ in range(10): inlined = ir @@ -142,7 +145,8 @@ def apply_common_transforms( ir = NormalizeShifts().visit(ir) - ir = FuseMaps(uids=uids).visit(ir) + if fuse_maps: + ir = FuseMaps(uids=uids).visit(ir) ir = CollapseListGet().visit(ir) if unroll_reduce: @@ -164,6 +168,12 @@ def apply_common_transforms( ir, opcount_preserving=True, force_inline_lambda_args=force_inline_lambda_args ) + ir = infer_domain.infer_program( + ir, + offset_provider=offset_provider, + symbolic_domain_sizes=symbolic_domain_sizes, + ) + assert isinstance(ir, itir.Program) return ir diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index da7f832fa3..722290b965 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -383,7 +383,12 @@ def _generate_sdfg_without_configuring_dace( column_axis: Optional[common.Dimension], ) -> dace.SDFG: if not self.disable_itir_transforms: - ir = itir_transforms.apply_fieldview_transforms(ir, offset_provider=offset_provider) + ir = itir_transforms.apply_common_transforms( + ir, + offset_provider=offset_provider, + fuse_maps=False, + transform_concat_where_to_as_fieldop=False, + ) offset_provider_type = common.offset_provider_to_type(offset_provider) on_gpu = self.device_type != core_defs.DeviceType.CPU From 782f5ff5c7fafd88900f3da6d4e383a067648d29 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Mon, 2 Feb 2026 17:13:29 +0100 Subject: [PATCH 02/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 74 ++++++++++++------- .../dace/lowering/gtir_to_sdfg_primitives.py | 49 +++++++----- 2 files changed, 78 insertions(+), 45 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 562da246ab..f9db02ad72 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -1527,7 +1527,6 @@ def _make_cartesian_shift( self, it: IteratorExpr, offset_dim: gtx_common.Dimension, offset_expr: DataExpr ) -> IteratorExpr: """Implements cartesian shift along one dimension.""" - assert any(dim == offset_dim for dim, _ in it.field_domain) new_index: SymbolExpr | ValueExpr index_expr = it.indices[offset_dim] if isinstance(index_expr, SymbolExpr) and isinstance(offset_expr, SymbolExpr): @@ -1599,9 +1598,9 @@ def _make_cartesian_shift( def _make_dynamic_neighbor_offset( self, - offset_expr: MemletExpr | ValueExpr, + offset_expr: MemletExpr | ValueExpr | SymbolExpr, offset_table_node: dace.nodes.AccessNode, - origin_index: SymbolExpr, + origin_index: MemletExpr | ValueExpr | SymbolExpr, ) -> ValueExpr: """ Implements access to neighbor connectivity table by means of a tasklet node. @@ -1610,33 +1609,56 @@ def _make_dynamic_neighbor_offset( or computed by another tasklet (`DataExpr`). """ new_index_connector = "neighbor_index" - tasklet_node, connector_mapping = self._add_tasklet( - "dynamic_neighbor_offset", - {"table", "offset"}, - {new_index_connector}, - f"{new_index_connector} = table[{origin_index.value}, offset]", - ) + if isinstance(offset_expr, SymbolExpr) and isinstance(origin_index, SymbolExpr): + tasklet_node, connector_mapping = self._add_tasklet( + "dynamic_neighbor_offset", + {"table"}, + {new_index_connector}, + f"{new_index_connector} = table[{origin_index.value}, {offset_expr.value}]", + ) + elif isinstance(origin_index, SymbolExpr): + tasklet_node, connector_mapping = self._add_tasklet( + "dynamic_neighbor_offset", + {"table", "offset"}, + {new_index_connector}, + f"{new_index_connector} = table[{origin_index.value}, offset]", + ) + elif isinstance(offset_expr, SymbolExpr): + tasklet_node, connector_mapping = self._add_tasklet( + "dynamic_neighbor_offset", + {"table", "origin"}, + {new_index_connector}, + f"{new_index_connector} = table[origin, {offset_expr.value}]", + ) + else: + tasklet_node, connector_mapping = self._add_tasklet( + "dynamic_neighbor_offset", + {"table", "offset", "origin"}, + {new_index_connector}, + f"{new_index_connector} = table[origin, offset]", + ) self._add_input_data_edge( offset_table_node, dace_subsets.Range.from_array(offset_table_node.desc(self.sdfg)), tasklet_node, connector_mapping["table"], ) - if isinstance(offset_expr, MemletExpr): - self._add_input_data_edge( - offset_expr.dc_node, - offset_expr.subset, - tasklet_node, - connector_mapping["offset"], - ) - else: - self._add_edge( - offset_expr.dc_node, - None, - tasklet_node, - connector_mapping["offset"], - dace.Memlet(data=offset_expr.dc_node.data, subset="0"), - ) + for conn, input_expr in [("offset", offset_expr), ("origin", origin_index)]: + if isinstance(input_expr, MemletExpr): + self._add_input_data_edge( + input_expr.dc_node, + input_expr.subset, + tasklet_node, + connector_mapping[conn], + ) + elif isinstance(input_expr, ValueExpr): + self._add_edge( + input_expr.dc_node, + None, + tasklet_node, + connector_mapping[conn], + dace.Memlet(data=input_expr.dc_node.data, subset="0"), + ) dc_dtype = offset_table_node.desc(self.sdfg).dtype return self._construct_tasklet_result( @@ -1651,14 +1673,12 @@ def _make_unstructured_shift( offset_expr: DataExpr, ) -> IteratorExpr: """Implements shift in unstructured domain by means of a neighbor table.""" - assert any(dim == conn_type.codomain for dim, _ in it.field_domain) neighbor_dim = conn_type.codomain origin_dim = conn_type.source_dim origin_index = it.indices[origin_dim] - assert isinstance(origin_index, SymbolExpr) shifted_indices = {dim: idx for dim, idx in it.indices.items() if dim != origin_dim} - if isinstance(offset_expr, SymbolExpr): + if isinstance(offset_expr, SymbolExpr) and isinstance(origin_index, SymbolExpr): # use memlet to retrieve the neighbor index shifted_indices[neighbor_dim] = MemletExpr( dc_node=conn_node, diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index 2428607236..0d6df88275 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -14,6 +14,7 @@ import dace from dace import subsets as dace_subsets +from gt4py.eve.extended_typing import MaybeNestedInTuple from gt4py.next import common as gtx_common, utils as gtx_utils from gt4py.next.iterator import ir as gtir from gt4py.next.iterator.ir_utils import ( @@ -73,16 +74,18 @@ def _parse_fieldop_arg( ctx: gtir_to_sdfg.SubgraphContext, sdfg_builder: gtir_to_sdfg.SDFGBuilder, domain: gtir_domain.FieldopDomain, -) -> gtir_dataflow.IteratorExpr | gtir_dataflow.MemletExpr: +) -> MaybeNestedInTuple[gtir_dataflow.IteratorExpr | gtir_dataflow.MemletExpr]: """ Helper method to visit an expression passed as argument to a field operator and create the local view for the field argument. """ arg = sdfg_builder.visit(node, ctx=ctx) - if not isinstance(arg, gtir_to_sdfg_types.FieldopData): - raise ValueError("Expected a field, found a tuple of fields.") - return arg.get_local_view(domain, ctx.sdfg) + if isinstance(arg, gtir_to_sdfg_types.FieldopData): + return arg.get_local_view(domain, ctx.sdfg) + else: + # handle tuples of fields + return gtx_utils.tree_map(lambda targ: targ.get_local_view(domain))(arg) def _create_field_operator_impl( @@ -175,10 +178,10 @@ def _create_field_operator_impl( def _create_field_operator( ctx: gtir_to_sdfg.SubgraphContext, domain: gtir_domain.FieldopDomain, - node_type: ts.FieldType, + node_type: ts.FieldType | ts.TupleType, sdfg_builder: gtir_to_sdfg.SDFGBuilder, input_edges: Iterable[gtir_dataflow.DataflowInputEdge], - output_edge: gtir_dataflow.DataflowOutputEdge, + output_tree: MaybeNestedInTuple[gtir_dataflow.DataflowOutputEdge], ) -> gtir_to_sdfg_types.FieldopResult: """ Helper method to build the output of a field operator. @@ -189,11 +192,11 @@ def _create_field_operator( node_type: The GT4Py type of the IR node that produces this field. sdfg_builder: The object used to build the map scope in the provided SDFG. input_edges: List of edges to pass input data into the dataflow. - output_edge: Edge corresponding to the dataflow output. + output_tree: A tree representation of the dataflow output data. Returns: - The descriptor of the field operator result, which is a single field defined - on the domain of the field operator. + The descriptor of the field operator result, which can be either a single + field or a tuple fields. """ if len(domain) == 0: @@ -215,7 +218,19 @@ def _create_field_operator( for edge in input_edges: edge.connect(map_entry) - return _create_field_operator_impl(ctx, sdfg_builder, domain, output_edge, node_type, map_exit) + if isinstance(node_type, ts.FieldType): + assert isinstance(output_tree, gtir_dataflow.DataflowOutputEdge) + return _create_field_operator_impl( + ctx, sdfg_builder, domain, output_tree, node_type, map_exit + ) + else: + # handle tuples of fields + output_symbol_tree = gtir_to_sdfg_utils.make_symbol_tree("x", node_type) + return gtx_utils.tree_map( + lambda output_edge, output_sym: _create_field_operator_impl( + ctx, sdfg_builder, domain, output_edge, output_sym.type, map_exit + ) + )(output_tree, output_symbol_tree) def translate_as_fieldop( @@ -247,9 +262,6 @@ def translate_as_fieldop( if cpm.is_call_to(fieldop_expr, "scan"): return translate_scan(node, ctx, sdfg_builder) - if not isinstance(node.type, ts.FieldType): - raise NotImplementedError("Unexpected 'as_fieldop' with tuple output in SDFG lowering.") - # Parse the domain of the field operator. assert isinstance(fieldop_domain_expr.type, ts.DomainType) field_domain = gtir_domain.get_field_domain( @@ -259,11 +271,13 @@ def translate_as_fieldop( if cpm.is_ref_to(fieldop_expr, "deref"): arg_type = node.args[0].type assert isinstance(arg_type, (ts.FieldType, ts.ScalarType)) - if isinstance(arg_type, ts.ScalarType) or arg_type.dims != node.type.dims: + if ( + isinstance(arg_type, ts.ScalarType) or arg_type.dims != node.type.dims # type: ignore[union-attr] + ): # Special usage of 'deref' as argument to fieldop expression, to broadcast # the input value (a scalar or a field slice) on the output domain. stencil_expr = im.lambda_("a")(im.deref("a")) - stencil_expr.expr.type = node.type.dtype + stencil_expr.expr.type = node.type.dtype # type: ignore[union-attr] else: # Special usage of 'deref' with field argument, to access the field # on the given domain. It copies a subset of the source field. @@ -283,13 +297,12 @@ def translate_as_fieldop( fieldop_args = [_parse_fieldop_arg(arg, ctx, sdfg_builder, field_domain) for arg in node.args] # represent the field operator as a mapped tasklet graph, which will range over the field domain - input_edges, output_edge = gtir_dataflow.translate_lambda_to_dataflow( + input_edges, output_edges = gtir_dataflow.translate_lambda_to_dataflow( ctx.sdfg, ctx.state, sdfg_builder, stencil_expr, fieldop_args ) - assert isinstance(output_edge, gtir_dataflow.DataflowOutputEdge) return _create_field_operator( - ctx, field_domain, node.type, sdfg_builder, input_edges, output_edge + ctx, field_domain, node.type, sdfg_builder, input_edges, output_edges ) From 2d24ba994e6208d77b7d2be104722cd5a18416b4 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 9 Apr 2026 21:51:26 +0200 Subject: [PATCH 03/49] enable warning --- .../next/iterator/transforms/pass_manager.py | 17 +++++++++-------- 1 file changed, 9 insertions(+), 8 deletions(-) diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index a6228c6125..045d3bcd3f 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -5,6 +5,7 @@ # # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import warnings from typing import Optional, Protocol from gt4py.next import common, utils @@ -113,15 +114,15 @@ def _process_symbolic_domains_option( if use_max_domain_range_on_unstructured_shift is None: use_max_domain_range_on_unstructured_shift = _has_dynamic_domains(ir) + elif use_max_domain_range_on_unstructured_shift: + if not _has_dynamic_domains(ir): + warnings.warn( + "You are using static domains together with " + "'use_max_domain_range_on_unstructured_shift'. This is " + "likely not what you wanted.", + stacklevel=2, + ) # noqa: ERA001, RUF100 if use_max_domain_range_on_unstructured_shift: - # TODO(havogt): ICON4Py uses this codepath as default for now. Once we use the minimal domain range, we should re-enable this warning. - # if not _has_dynamic_domains(ir): - # warnings.warn( - # "You are using static domains together with " - # "'use_max_domain_range_on_unstructured_shift'. This is " - # "likely not what you wanted.", - # stacklevel=2, # noqa: ERA001 - # ) # noqa: ERA001, RUF100 assert not symbolic_domain_sizes, "Options are mutually exclusive." symbolic_domain_sizes = _max_domain_range_sizes(offset_provider) # type: ignore[assignment] return symbolic_domain_sizes From e3ac9de64b9b5fde96c53da4f91b02c6065b591b Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Mon, 13 Apr 2026 14:51:56 +0200 Subject: [PATCH 04/49] edit --- .../runners/dace/workflow/translation.py | 36 ++++++++++++++----- 1 file changed, 28 insertions(+), 8 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index 74ffa1fa09..a4340d0b97 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -9,6 +9,7 @@ from __future__ import annotations import dataclasses +import functools from typing import Any, Optional import dace @@ -17,7 +18,8 @@ from gt4py._core import definitions as core_defs from gt4py.next import common, config from gt4py.next.instrumentation import metrics -from gt4py.next.iterator import ir as itir, transforms as itir_transforms +from gt4py.next.iterator import ir as itir +from gt4py.next.iterator.transforms import pass_manager from gt4py.next.otf import code_specs, definitions, stages, workflow from gt4py.next.otf.binding import interface from gt4py.next.program_processors.runners.dace import ( @@ -375,6 +377,30 @@ def generate_sdfg( with gtx_wfdcommon.dace_context(device_type=self.device_type): return self._generate_sdfg_without_configuring_dace(*args, **kwargs) + def _preprocess_program( + self, + program: itir.Program, + offset_provider: common.OffsetProvider | common.OffsetProviderType, + ) -> itir.Program: + apply_common_transforms = functools.partial( + pass_manager.apply_common_transforms, + offset_provider=offset_provider, + transform_concat_where_to_as_fieldop=False, + use_max_domain_range_on_unstructured_shift=self.use_max_domain_range_on_unstructured_shift, + ) + + new_program = apply_common_transforms(program, unroll_reduce=False) + + if any( + node.id == "neighbors" + for node in new_program.pre_walk_values().if_isinstance(itir.SymRef) + ): + # if we don't unroll, there may be lifts left in the itir which can't + # be lowered to SDFG. In this case, just retry with unrolled reductions. + new_program = apply_common_transforms(program, unroll_reduce=True) + + return new_program + def _generate_sdfg_without_configuring_dace( self, ir: itir.Program, @@ -382,13 +408,7 @@ def _generate_sdfg_without_configuring_dace( column_axis: Optional[common.Dimension], ) -> dace.SDFG: if not self.disable_itir_transforms: - ir = itir_transforms.apply_common_transforms( - ir, - offset_provider=offset_provider, - fuse_maps=False, - transform_concat_where_to_as_fieldop=False, - use_max_domain_range_on_unstructured_shift=self.use_max_domain_range_on_unstructured_shift, - ) + ir = self._preprocess_program(ir, offset_provider) offset_provider_type = common.offset_provider_to_type(offset_provider) on_gpu = self.device_type != core_defs.DeviceType.CPU From d33bd83b7083b0f3f3466cfc9d4b4c6c5bcb23a1 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 15 Apr 2026 10:28:01 +0200 Subject: [PATCH 05/49] fix sdfg lowering --- .../runners/dace/lowering/gtir_dataflow.py | 44 ++++++++++++------- 1 file changed, 29 insertions(+), 15 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index da590d84e0..e7dc855428 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -989,11 +989,18 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: it = self.visit(node.args[1]) assert isinstance(it, IteratorExpr) + if not all(isinstance(index, SymbolExpr) for index in it.indices.values()): + raise NotImplementedError("Dynamic indices in neighbors expression are not supported.") + + # make sure that the field can be dereferenced with the given connectivity type assert any(dim == conn_type.codomain for dim, _ in it.field_domain) + field_codomain_origin = next( + origin for dim, origin in it.field_domain if dim == conn_type.codomain + ) + # make sure that the iterator can access the connectivity table assert conn_type.source_dim in it.indices - origin_index = it.indices[conn_type.source_dim] - assert isinstance(origin_index, SymbolExpr) - assert all(isinstance(index, SymbolExpr) for index in it.indices.values()) + conn_source_index = it.indices[conn_type.source_dim] + assert isinstance(conn_source_index, SymbolExpr) # initially, the storage for the connectivty tables is created as transient; # when the tables are used, the storage is changed to non-transient, @@ -1040,7 +1047,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: ), ), subset=dace_subsets.Range.from_string( - f"{origin_index.value}, 0:{conn_type.max_neighbors}" + f"{conn_source_index.value}, 0:{conn_type.max_neighbors}" ), ) ) @@ -1055,7 +1062,9 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: index_connector = "__index" field_connector = "__field" output_connector = "__val" - tasklet_expression = f"{output_connector} = {field_connector}[{index_connector}]" + tasklet_expression = ( + f"{output_connector} = {field_connector}[{index_connector} - {field_codomain_origin}]" + ) input_memlets = { field_connector: self.sdfg.make_array_memlet(field_slice.dc_node.data), index_connector: dace.Memlet(data=conn_slice.dc_node.data, subset=neighbor_idx), @@ -1651,29 +1660,34 @@ def _make_unstructured_shift( offset_expr: DataExpr, ) -> IteratorExpr: """Implements shift in unstructured domain by means of a neighbor table.""" + # make sure that the field can be dereferenced with the given connectivity type assert any(dim == conn_type.codomain for dim, _ in it.field_domain) - neighbor_dim = conn_type.codomain - origin_dim = conn_type.source_dim - origin_index = it.indices[origin_dim] - assert isinstance(origin_index, SymbolExpr) + # make sure that the iterator can access the connectivity table + assert conn_type.source_dim in it.indices + conn_source_index = it.indices[conn_type.source_dim] + assert isinstance(conn_source_index, SymbolExpr) - shifted_indices = {dim: idx for dim, idx in it.indices.items() if dim != origin_dim} + shifted_indices = { + dim: idx for dim, idx in it.indices.items() if dim != conn_type.source_dim + } if isinstance(offset_expr, SymbolExpr): # use memlet to retrieve the neighbor index - shifted_indices[neighbor_dim] = MemletExpr( + shifted_indices[conn_type.codomain] = MemletExpr( dc_node=conn_node, gt_field=ts.FieldType( - dims=[origin_dim], + dims=[conn_type.source_dim], dtype=ts.ListType( element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM ), ), - subset=dace_subsets.Range.from_string(f"{origin_index.value}, {offset_expr.value}"), + subset=dace_subsets.Range.from_string( + f"{conn_source_index.value}, {offset_expr.value}" + ), ) else: # dynamic offset: we cannot use a memlet to retrieve the offset value, use a tasklet node - shifted_indices[neighbor_dim] = self._make_dynamic_neighbor_offset( - offset_expr, conn_node, origin_index + shifted_indices[conn_type.codomain] = self._make_dynamic_neighbor_offset( + offset_expr, conn_node, conn_source_index ) return IteratorExpr(it.field, it.gt_dtype, it.field_domain, shifted_indices) From 3d770bd3c407d0e6ede665ba4a32134ece336fb3 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 15 Apr 2026 12:05:33 +0200 Subject: [PATCH 06/49] fix bindings test for unstructured grid --- .../dace_tests/test_dace_bindings.py | 25 ++++++++++--------- 1 file changed, 13 insertions(+), 12 deletions(-) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index 855b315009..25dae344f2 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -157,7 +157,7 @@ def {_bind_func_name}(device, sdfg_argtypes, args, sdfg_call_args, offset_provid def _binding_source_unstructured(use_metrics: bool) -> str: metrics_arg_index = 2 - idx = [0, 4, 1, 5, 6, 7, 2, 9, 8, 3, 11, 10] + idx = [0, 4, 5, 1, 6, 7, 8, 2, 10, 9, 3, 12, 11] if use_metrics: idx = [idx + 1 if idx >= metrics_arg_index else idx for idx in idx] return ( @@ -169,19 +169,20 @@ def {_bind_func_name}(device, sdfg_argtypes, args, sdfg_call_args, offset_provid args_1, ) = args sdfg_call_args[{idx[0]}].value = args_0.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[1]}] = ctypes.c_int(args_0.__gt_buffer_info__.elem_strides[0]) - sdfg_call_args[{idx[2]}].value = args_1.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[3]}] = ctypes.c_int(args_1.domain.ranges[0].start) - sdfg_call_args[{idx[4]}] = ctypes.c_int(args_1.domain.ranges[0].stop) - sdfg_call_args[{idx[5]}] = ctypes.c_int(args_1.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[1]}] = ctypes.c_int(args_0.domain.ranges[0].start) + sdfg_call_args[{idx[2]}] = ctypes.c_int(args_0.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[3]}].value = args_1.__gt_buffer_info__.data_ptr + sdfg_call_args[{idx[4]}] = ctypes.c_int(args_1.domain.ranges[0].start) + sdfg_call_args[{idx[5]}] = ctypes.c_int(args_1.domain.ranges[0].stop) + sdfg_call_args[{idx[6]}] = ctypes.c_int(args_1.__gt_buffer_info__.elem_strides[0]) table_E2V = offset_provider["E2V"] - sdfg_call_args[{idx[6]}].value = table_E2V.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[7]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[0]) - sdfg_call_args[{idx[8]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[1]) + sdfg_call_args[{idx[7]}].value = table_E2V.__gt_buffer_info__.data_ptr + sdfg_call_args[{idx[8]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[9]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[1]) table_V2E = offset_provider["V2E"] - sdfg_call_args[{idx[9]}].value = table_V2E.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[10]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[0]) - sdfg_call_args[{idx[11]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[1]) + sdfg_call_args[{idx[10]}].value = table_V2E.__gt_buffer_info__.data_ptr + sdfg_call_args[{idx[11]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[12]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[1]) """ ) From 7a70de904174c63957211313698e55432c2960dd Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 16 Apr 2026 09:04:02 +0200 Subject: [PATCH 07/49] remove noqa comment --- src/gt4py/next/iterator/transforms/pass_manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index 045d3bcd3f..1236a3209a 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -121,7 +121,7 @@ def _process_symbolic_domains_option( "'use_max_domain_range_on_unstructured_shift'. This is " "likely not what you wanted.", stacklevel=2, - ) # noqa: ERA001, RUF100 + ) if use_max_domain_range_on_unstructured_shift: assert not symbolic_domain_sizes, "Options are mutually exclusive." symbolic_domain_sizes = _max_domain_range_sizes(offset_provider) # type: ignore[assignment] From a36755302741d4f9304afad8125cf23706dfb31f Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 16 Apr 2026 09:43:03 +0200 Subject: [PATCH 08/49] add test coverage --- .../ffront_tests/test_execution.py | 40 +++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py index c58ac5f497..664229e7e8 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py @@ -140,6 +140,24 @@ def testee(a: cases.VField) -> cases.EField: ) +@pytest.mark.uses_unstructured_shift +def test_unstructured_shift_with_non_zero_origin(unstructured_case): + @gtx.field_operator + def testee(a: cases.VField) -> cases.EField: + return a(E2V[0]) + + a = cases.allocate(unstructured_case, testee, "a")() + out = cases.allocate(unstructured_case, testee, cases.RETURN)() + + ORIGIN = 2 + neighbor_0_iter = iter(enumerate(unstructured_case.offset_provider["E2V"].asnumpy()[:, 0])) + edge_start = next(i for i, v in neighbor_0_iter if v >= ORIGIN) + edge_stop = next(i for i, v in neighbor_0_iter if v < ORIGIN) + + ref = a.ndarray[unstructured_case.offset_provider["E2V"].asnumpy()[edge_start:edge_stop, 0]] + cases.verify(unstructured_case, testee, a[ORIGIN:], out=out[edge_start:edge_stop], ref=ref) + + def test_horizontal_only_with_3d_mesh(unstructured_case_3d): # test field operator operating only on horizontal fields while using an offset provider # including a vertical dimension. @@ -724,6 +742,28 @@ def combine(a: cases.IField, b: cases.IField) -> cases.IField: cases.verify_with_default_data(cartesian_case, combine, ref=lambda a, b: a + a + b) +@pytest.mark.uses_unstructured_shift +def test_neighbor_sum_with_non_zero_origin(unstructured_case): + @gtx.field_operator + def testee(a: cases.VField) -> cases.EField: + return neighbor_sum(a(E2V), axis=E2VDim) + + a = cases.allocate(unstructured_case, testee, "a")() + out = cases.allocate(unstructured_case, testee, cases.RETURN)() + + ORIGIN = 2 + neighbor_iter = iter(enumerate(unstructured_case.offset_provider["E2V"].asnumpy())) + edge_start = next(i for i, v in neighbor_iter if all(v >= ORIGIN)) + edge_stop = next(i for i, v in neighbor_iter if any(v < ORIGIN)) + + ref = np.sum( + a.ndarray[unstructured_case.offset_provider["E2V"].asnumpy()[edge_start:edge_stop,]], + axis=1, + initial=0.0, + ) + cases.verify(unstructured_case, testee, a[ORIGIN:], out=out[edge_start:edge_stop], ref=ref) + + @pytest.mark.uses_unstructured_shift def test_nested_reduction(unstructured_case): @gtx.field_operator From 33d60b7904b39bd84bd2b05139f81fb9c74e4741 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 16 Apr 2026 10:06:17 +0200 Subject: [PATCH 09/49] add test marker for sliced out argument --- .../feature_tests/ffront_tests/test_execution.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py index 664229e7e8..3b99af3b59 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py @@ -141,6 +141,7 @@ def testee(a: cases.VField) -> cases.EField: @pytest.mark.uses_unstructured_shift +@pytest.mark.uses_program_with_sliced_out_arguments def test_unstructured_shift_with_non_zero_origin(unstructured_case): @gtx.field_operator def testee(a: cases.VField) -> cases.EField: @@ -743,6 +744,7 @@ def combine(a: cases.IField, b: cases.IField) -> cases.IField: @pytest.mark.uses_unstructured_shift +@pytest.mark.uses_program_with_sliced_out_arguments def test_neighbor_sum_with_non_zero_origin(unstructured_case): @gtx.field_operator def testee(a: cases.VField) -> cases.EField: From 75e062ace91b064a6331a99a44ced703f0947a83 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 16 Apr 2026 10:59:31 +0200 Subject: [PATCH 10/49] remove initial from np.sum() --- .../feature_tests/ffront_tests/test_execution.py | 14 ++++++-------- 1 file changed, 6 insertions(+), 8 deletions(-) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py index 3b99af3b59..dcb2d73da1 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py @@ -151,11 +151,12 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - neighbor_0_iter = iter(enumerate(unstructured_case.offset_provider["E2V"].asnumpy()[:, 0])) + e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + neighbor_0_iter = iter(enumerate(e2v_table[:, 0])) edge_start = next(i for i, v in neighbor_0_iter if v >= ORIGIN) edge_stop = next(i for i, v in neighbor_0_iter if v < ORIGIN) - ref = a.ndarray[unstructured_case.offset_provider["E2V"].asnumpy()[edge_start:edge_stop, 0]] + ref = a.ndarray[e2v_table[edge_start:edge_stop, 0]] cases.verify(unstructured_case, testee, a[ORIGIN:], out=out[edge_start:edge_stop], ref=ref) @@ -754,15 +755,12 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - neighbor_iter = iter(enumerate(unstructured_case.offset_provider["E2V"].asnumpy())) + e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + neighbor_iter = iter(enumerate(e2v_table)) edge_start = next(i for i, v in neighbor_iter if all(v >= ORIGIN)) edge_stop = next(i for i, v in neighbor_iter if any(v < ORIGIN)) - ref = np.sum( - a.ndarray[unstructured_case.offset_provider["E2V"].asnumpy()[edge_start:edge_stop,]], - axis=1, - initial=0.0, - ) + ref = np.sum(a.ndarray[e2v_table[edge_start:edge_stop,]], axis=1) cases.verify(unstructured_case, testee, a[ORIGIN:], out=out[edge_start:edge_stop], ref=ref) From a22ff35fdf99aac90b7d3bd46e857abfb9f12db1 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Tue, 5 May 2026 11:16:48 +0200 Subject: [PATCH 11/49] add xfail for embedded backend --- .../feature_tests/ffront_tests/test_execution.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py index dcb2d73da1..af1410b2be 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_execution.py @@ -143,6 +143,9 @@ def testee(a: cases.VField) -> cases.EField: @pytest.mark.uses_unstructured_shift @pytest.mark.uses_program_with_sliced_out_arguments def test_unstructured_shift_with_non_zero_origin(unstructured_case): + if unstructured_case.backend is None: + pytest.xfail("Embedded backend requires contiguous inverse image.") + @gtx.field_operator def testee(a: cases.VField) -> cases.EField: return a(E2V[0]) @@ -747,6 +750,9 @@ def combine(a: cases.IField, b: cases.IField) -> cases.IField: @pytest.mark.uses_unstructured_shift @pytest.mark.uses_program_with_sliced_out_arguments def test_neighbor_sum_with_non_zero_origin(unstructured_case): + if unstructured_case.backend is None: + pytest.xfail("Embedded backend requires contiguous inverse image.") + @gtx.field_operator def testee(a: cases.VField) -> cases.EField: return neighbor_sum(a(E2V), axis=E2VDim) From 0d17ef643f372272120b9ed001c7e9fda253ed40 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Tue, 5 May 2026 15:07:42 +0200 Subject: [PATCH 12/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 96 +++++++++++++++++-- .../dace/lowering/gtir_to_sdfg_primitives.py | 2 +- .../dace/lowering/gtir_to_sdfg_scan.py | 38 ++------ .../runners/dace/workflow/translation.py | 1 + 4 files changed, 100 insertions(+), 37 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 8d1678274d..91985530a8 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -552,6 +552,51 @@ def _construct_tasklet_result( ), ) + def _visit_can_deref(self, node: gtir.FunCall) -> DataExpr: + assert isinstance(node.type, ts.ScalarType) and node.type.kind == ts.ScalarKind.BOOL + assert len(node.args) == 1 + if not cpm.is_applied_shift(node.args[0]): + raise NotImplementedError( + f"Only `can_deref` of unstructured `shift` expressions is supported, got {node.args[0]}." + ) + it = self._visit_shift(node.args[0]) + index_values = {k: v for k, v in it.indices.items() if not isinstance(v, SymbolExpr)} + if len(index_values) == 0: + raise ValueError(f"Unexpected `can_deref` argument: {it}.") + can_deref_node, connector_mapping = self._add_tasklet( + name="can_deref", + inputs={f"index_{dim.value}" for dim in index_values}, + outputs={"valid"}, + code="valid = " + + " and ".join( + f"index_{dim.value} != {gtx_common._DEFAULT_SKIP_VALUE}" + for dim in index_values.keys() + ), + ) + for dim, index_expr in index_values.items(): + index_connector = f"index_{dim.value}" + if isinstance(index_expr, MemletExpr): + self._add_input_data_edge( + index_expr.dc_node, + index_expr.subset, + can_deref_node, + connector_mapping[index_connector], + ) + + else: + self._add_edge( + index_expr.dc_node, + None, + can_deref_node, + connector_mapping[index_connector], + dace.Memlet(data=index_expr.dc_node.data, subset="0"), + ) + return self._construct_tasklet_result( + dc_dtype=dace.bool_, + src_node=can_deref_node, + src_connector=connector_mapping["valid"], + ) + def _visit_deref(self, node: gtir.FunCall) -> DataExpr: """ Visit a `deref` node, which represents dereferencing of an iterator. @@ -759,10 +804,15 @@ def _visit_if_branch( """ assert if_branch_state in if_sdfg.states() - lambda_args = [] - lambda_params = [] + lambda_args: list[MaybeNestedInTuple[IteratorExpr | DataExpr]] = [] + lambda_params: list[gtir.Sym] = [] for pname in symbol_ref_utils.collect_symbol_refs(expr, self.symbol_map.keys()): arg = self.symbol_map[pname] + if isinstance(arg, SymbolExpr): + psymbol = im.sym(pname, gtx_dace_args.as_itir_type(arg.dc_dtype)) + lambda_args.append(arg) + lambda_params.append(psymbol) + continue if isinstance(arg, tuple): ptype = get_tuple_type(arg) # type: ignore[arg-type] psymbol = im.sym(pname, ptype) @@ -781,7 +831,7 @@ def _visit_if_branch( ) )(psymbol_tree, arg) else: - psymbol = im.sym(pname, arg.gt_dtype) # type: ignore[union-attr] + psymbol = im.sym(pname, arg.gt_dtype) deref_on_input_memlet = pname in direct_deref_iterators inner_arg = self._visit_if_branch_arg( if_sdfg, @@ -868,6 +918,13 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp nsdfg = dace.SDFG(self.subgraph_builder.unique_nsdfg_name("if_stmt")) nsdfg.debuginfo = gtir_to_sdfg_utils.debug_info(node, default=self.sdfg.debuginfo) + # add connectivities + for aname, adesc in self.sdfg.arrays.items(): + if gtx_dace_args.is_connectivity_identifier(aname): + adesc = adesc.clone() + adesc.transient = True + nsdfg.add_datadesc(aname, adesc) + # create states inside the nested SDFG for the if-branches if_region = dace.sdfg.state.ConditionalBlock("if") nsdfg.add_node(if_region, ensure_unique_name=True) @@ -948,13 +1005,21 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp outputs = {outval.dc_node.data for outval in gtx_utils.flatten_nested_tuple((result,))} + # map the connectivities that were used inside the nested SDFG + used_connectivities = { + aname + for aname, adesc in nsdfg.arrays.items() + if gtx_dace_args.is_connectivity_identifier(aname) and not adesc.transient + } + # all free symbols are mapped to the symbols available in parent SDFG nsdfg_symbols_mapping = {str(sym): sym for sym in nsdfg.free_symbols} if isinstance(condition_value, SymbolExpr): nsdfg_symbols_mapping["__cond"] = condition_value.value + nsdfg_node = self.state.add_nested_sdfg( nsdfg, - inputs=set(input_memlets.keys()), + inputs=(used_connectivities | input_memlets.keys()), outputs=outputs, symbol_mapping=nsdfg_symbols_mapping, ) @@ -971,6 +1036,16 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp self.sdfg.make_array_memlet(input_expr.dc_node.data), ) + for conn in used_connectivities: + desc = self.sdfg.data(conn) + desc.transient = False + self._add_input_data_edge( + self.state.add_access(conn), + dace_subsets.Range.from_array(desc), + nsdfg_node, + conn, + ) + return ( gtx_utils.tree_map(write_output_of_nested_sdfg_to_temporary)(result) if isinstance(result, tuple) @@ -1122,7 +1197,7 @@ def _visit_list_get(self, node: gtir.FunCall) -> ValueExpr: assert index_arg.dc_dtype in dace.dtypes.INTEGER_TYPES src_subset = ( dace_subsets.Range(src_subset[:local_dim_index]) - + dace_subsets.Range.from_string(index_arg.value) + + dace_subsets.Range.from_indices([index_arg.value]) + dace_subsets.Range(src_subset[local_dim_index + 1 :]) ) if isinstance(src_arg, MemletExpr): @@ -1825,7 +1900,10 @@ def _visit_tuple_get( return tuple_fields[index] def visit_FunCall(self, node: gtir.FunCall) -> MaybeNestedInTuple[IteratorExpr | DataExpr]: - if cpm.is_call_to(node, "deref"): + if cpm.is_call_to(node, "can_deref"): + return self._visit_can_deref(node) + + elif cpm.is_call_to(node, "deref"): return self._visit_deref(node) elif cpm.is_call_to(node, "if_"): @@ -1910,6 +1988,9 @@ def visit_Literal(self, node: gtir.Literal) -> SymbolExpr: dc_dtype = gtx_dace_args.as_dace_type(node.type) return SymbolExpr(node.value, dc_dtype) + def visit_OffsetLiteral(self, node: gtir.OffsetLiteral) -> SymbolExpr: + return SymbolExpr(node.value, gtir_to_sdfg_types.INDEX_DTYPE) + def visit_SymRef(self, node: gtir.SymRef) -> MaybeNestedInTuple[IteratorExpr | DataExpr]: param = str(node.id) if param in self.symbol_map: @@ -1924,7 +2005,7 @@ def translate_lambda_to_dataflow( state: dace.SDFGState, sdfg_builder: gtir_to_sdfg.DataflowBuilder, node: gtir.Lambda, - args: Sequence[MaybeNestedInTuple[IteratorExpr | MemletExpr | ValueExpr]], + args: Sequence[MaybeNestedInTuple[IteratorExpr | DataExpr]], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ Entry point to visit a `Lambda` node and lower it to a dataflow graph, @@ -1956,6 +2037,7 @@ def translate_lambda_to_dataflow( flat_arg_nodes = ( x.field if isinstance(x, IteratorExpr) else x.dc_node # type: ignore[attr-defined] for x in gtx_utils.flatten_nested_tuple(tuple(args)) + if not isinstance(x, SymbolExpr) ) state.remove_nodes_from([node for node in flat_arg_nodes if state.degree(node) == 0]) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index e04bcc13a5..3b5cff5ad4 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -85,7 +85,7 @@ def _parse_fieldop_arg( return arg.get_local_view(domain, ctx.sdfg) else: # handle tuples of fields - return gtx_utils.tree_map(lambda targ: targ.get_local_view(domain))(arg) + return gtx_utils.tree_map(lambda targ: targ.get_local_view(domain, ctx.sdfg))(arg) def _create_field_operator_impl( diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index a569c06fbf..f315ce221f 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -47,43 +47,23 @@ from gt4py.next.type_system import type_info as ti, type_specifications as ts -def _parse_scan_fieldop_arg( +def _parse_fieldop_arg( node: gtir.Expr, ctx: gtir_to_sdfg.SubgraphContext, sdfg_builder: gtir_to_sdfg.SDFGBuilder, - field_domain: gtir_domain.FieldopDomain, -) -> MaybeNestedInTuple[gtir_dataflow.MemletExpr]: - """Helper method to visit an expression passed as argument to a scan field operator. - - On the innermost level, a scan operator is lowered to a loop region which computes - column elements in the vertical dimension. - - It differs from the helper method `gtir_to_sdfg_primitives` in that field arguments - are passed in full shape along the vertical dimension, rather than as iterator. + domain: gtir_domain.FieldopDomain, +) -> MaybeNestedInTuple[gtir_dataflow.IteratorExpr | gtir_dataflow.MemletExpr]: + """ + Helper method to visit an expression passed as argument to a field operator + and create the local view for the field argument. """ - - def _parse_fieldop_arg_impl( - arg: gtir_to_sdfg_types.FieldopData, - ) -> gtir_dataflow.MemletExpr: - arg_expr = arg.get_local_view(field_domain, ctx.sdfg) - if isinstance(arg_expr, gtir_dataflow.MemletExpr): - return arg_expr - # In scan field operator, the arguments to the vertical stencil are passed by value. - # Therefore, the full field shape is passed as `MemletExpr` rather than `IteratorExpr`. - field_type = ts.FieldType( - dims=[dim for dim, _ in arg_expr.field_domain], dtype=arg_expr.gt_dtype - ) - return gtir_dataflow.MemletExpr( - arg_expr.field, field_type, arg_expr.get_memlet_subset(ctx.sdfg) - ) - arg = sdfg_builder.visit(node, ctx=ctx) if isinstance(arg, gtir_to_sdfg_types.FieldopData): - return _parse_fieldop_arg_impl(arg) + return arg.get_local_view(domain, ctx.sdfg) else: # handle tuples of fields - return gtx_utils.tree_map(_parse_fieldop_arg_impl)(arg) + return gtx_utils.tree_map(lambda targ: targ.get_local_view(domain, ctx.sdfg))(arg) def _create_scan_field_operator_impl( @@ -427,7 +407,7 @@ def get_scan_output_shape( # inside the 'compute' state, visit the list of arguments to be passed to the stencil stencil_args = [ - _parse_scan_fieldop_arg(im.ref(p.id), compute_ctx, sdfg_builder, field_domain) + _parse_fieldop_arg(im.ref(p.id), compute_ctx, sdfg_builder, field_domain) for p in lambda_node.params ] # stil inside the 'compute' state, generate the dataflow representing the stencil diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index a4340d0b97..820a895bf1 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -385,6 +385,7 @@ def _preprocess_program( apply_common_transforms = functools.partial( pass_manager.apply_common_transforms, offset_provider=offset_provider, + force_inline_lambda_args=True, transform_concat_where_to_as_fieldop=False, use_max_domain_range_on_unstructured_shift=self.use_max_domain_range_on_unstructured_shift, ) From a6d00a9c793636923eec504077686deced7d89f1 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Tue, 5 May 2026 15:35:20 +0200 Subject: [PATCH 13/49] add backends to test matrix --- .../runners/dace/__init__.py | 4 ++++ .../runners/dace/lowering/gtir_dataflow.py | 2 -- .../runners/dace/workflow/backend.py | 19 +++++++++++++++++++ .../runners/dace/workflow/translation.py | 10 +++++++++- tests/next_tests/definitions.py | 4 ++++ .../ffront_tests/ffront_test_utils.py | 8 ++++++++ tests/next_tests/unit_tests/conftest.py | 4 ++++ 7 files changed, 48 insertions(+), 3 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/__init__.py b/src/gt4py/next/program_processors/runners/dace/__init__.py index 0bb2c40dc3..0406d85f25 100644 --- a/src/gt4py/next/program_processors/runners/dace/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/__init__.py @@ -12,6 +12,8 @@ make_dace_backend, run_dace_cpu, run_dace_cpu_cached, + run_dace_cpu_gt, + run_dace_cpu_gt_noopt, run_dace_cpu_noopt, run_dace_gpu, run_dace_gpu_cached, @@ -24,6 +26,8 @@ "make_dace_backend", "run_dace_cpu", "run_dace_cpu_cached", + "run_dace_cpu_gt", + "run_dace_cpu_gt_noopt", "run_dace_cpu_noopt", "run_dace_gpu", "run_dace_gpu_cached", diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 4f8cd1547b..481c7b6138 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -1757,8 +1757,6 @@ def _make_unstructured_shift( offset_expr: DataExpr, ) -> IteratorExpr: """Implements shift in unstructured domain by means of a neighbor table.""" - # make sure that the field can be dereferenced with the given connectivity type - assert any(dim == conn_type.codomain for dim, _ in it.field_domain) # make sure that the iterator can access the connectivity table assert conn_type.source_dim in it.indices conn_source_index = it.indices[conn_type.source_dim] diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py index de6778a750..44be154f3c 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py @@ -69,6 +69,7 @@ class Params: def make_dace_backend( gpu: bool, cached: bool = True, + apply_common_transform: bool = False, auto_optimize: bool = True, async_sdfg_call: bool = True, optimization_args: dict[str, Any] | None = None, @@ -82,6 +83,8 @@ def make_dace_backend( Args: gpu: Enable GPU transformations and code generation. cached: Cache the lowered SDFG as a JSON file and the compiled programs. + apply_common_transform: Whether to apply the GTIR common transform before + lowering to SDFG. auto_optimize: Enable the SDFG auto-optimize pipeline. async_sdfg_call: Make an asynchronous SDFG call on GPU to allow overlapping of GPU kernel execution with the Python driver code. @@ -128,6 +131,7 @@ def make_dace_backend( cached=cached, auto_optimize=auto_optimize, otf_workflow__cached_translation=cached, + otf_workflow__bare_translation__apply_common_transform=apply_common_transform, otf_workflow__bare_translation__async_sdfg_call=(async_sdfg_call if gpu else False), otf_workflow__bare_translation__auto_optimize_args=optimization_args, otf_workflow__bare_translation__unstructured_horizontal_has_unit_stride=unstructured_horizontal_has_unit_stride, @@ -174,3 +178,18 @@ def make_dace_backend( auto_optimize=True, async_sdfg_call=True, ) + +run_dace_cpu_gt = make_dace_backend( + gpu=False, + cached=False, + apply_common_transform=True, + auto_optimize=True, + async_sdfg_call=False, +) +run_dace_cpu_gt_noopt = make_dace_backend( + gpu=False, + cached=False, + apply_common_transform=True, + auto_optimize=False, + async_sdfg_call=False, +) diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index 820a895bf1..7bf2482a0a 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -359,6 +359,7 @@ class DaCeTranslator( definitions.TranslationStep[code_specs.SDFGCodeSpec], ): device_type: core_defs.DeviceType + apply_common_transform: bool auto_optimize: bool auto_optimize_args: dict[str, Any] | None async_sdfg_call: bool @@ -409,7 +410,14 @@ def _generate_sdfg_without_configuring_dace( column_axis: Optional[common.Dimension], ) -> dace.SDFG: if not self.disable_itir_transforms: - ir = self._preprocess_program(ir, offset_provider) + if self.apply_common_transform: + ir = self._preprocess_program(ir, offset_provider) + else: + ir = pass_manager.apply_fieldview_transforms( + ir, + use_max_domain_range_on_unstructured_shift=self.use_max_domain_range_on_unstructured_shift, + offset_provider=offset_provider, + ) offset_provider_type = common.offset_provider_to_type(offset_provider) on_gpu = self.device_type != core_defs.DeviceType.CPU diff --git a/tests/next_tests/definitions.py b/tests/next_tests/definitions.py index 3396d93d3c..b440bae543 100644 --- a/tests/next_tests/definitions.py +++ b/tests/next_tests/definitions.py @@ -73,6 +73,8 @@ class OptionalProgramBackendId(_PythonObjectIdMixin, str, enum.Enum): DACE_CPU = "gt4py.next.program_processors.runners.dace.run_dace_cpu" DACE_GPU = "gt4py.next.program_processors.runners.dace.run_dace_gpu" DACE_CPU_NO_OPT = "gt4py.next.program_processors.runners.dace.run_dace_cpu_noopt" + DACE_CPU_GT = "gt4py.next.program_processors.runners.dace.run_dace_cpu_gt" + DACE_CPU_GT_NO_OPT = "gt4py.next.program_processors.runners.dace.run_dace_cpu_gt_noopt" class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): @@ -203,6 +205,8 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): OptionalProgramBackendId.DACE_CPU: DACE_SKIP_TEST_LIST, OptionalProgramBackendId.DACE_GPU: DACE_SKIP_TEST_LIST, OptionalProgramBackendId.DACE_CPU_NO_OPT: DACE_SKIP_TEST_LIST, + OptionalProgramBackendId.DACE_CPU_GT: DACE_SKIP_TEST_LIST, + OptionalProgramBackendId.DACE_CPU_GT_NO_OPT: DACE_SKIP_TEST_LIST, ProgramBackendId.GTFN_CPU: GTFN_SKIP_TEST_LIST + [(USES_SCAN_NESTED, XFAIL, UNSUPPORTED_MESSAGE)], ProgramBackendId.GTFN_CPU_IMPERATIVE: GTFN_SKIP_TEST_LIST diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/ffront_test_utils.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/ffront_test_utils.py index ab880868ce..ff129f1298 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/ffront_test_utils.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/ffront_test_utils.py @@ -103,6 +103,14 @@ def __gt_allocator__( next_tests.definitions.OptionalProgramBackendId.DACE_CPU_NO_OPT, marks=pytest.mark.requires_dace, ), + pytest.param( + next_tests.definitions.OptionalProgramBackendId.DACE_CPU_GT, + marks=pytest.mark.requires_dace, + ), + pytest.param( + next_tests.definitions.OptionalProgramBackendId.DACE_CPU_GT_NO_OPT, + marks=pytest.mark.requires_dace, + ), ], ids=lambda p: p.short_id(), ) diff --git a/tests/next_tests/unit_tests/conftest.py b/tests/next_tests/unit_tests/conftest.py index 580df0ee49..98fdb16ff8 100644 --- a/tests/next_tests/unit_tests/conftest.py +++ b/tests/next_tests/unit_tests/conftest.py @@ -65,6 +65,10 @@ def _program_processor(request) -> tuple[ProgramProcessor, bool]: (next_tests.definitions.OptionalProgramBackendId.DACE_CPU_NO_OPT, True), marks=pytest.mark.requires_dace, ), + pytest.param( + (next_tests.definitions.OptionalProgramBackendId.DACE_CPU_GT_NO_OPT, True), + marks=pytest.mark.requires_dace, + ), ], ids=lambda p: p[0].short_id() if p[0] is not None else "None", ) From 55f26f24794704dd232e7d3a5d6d0c04d1376432 Mon Sep 17 00:00:00 2001 From: "Philip Mueller, CSCS" Date: Wed, 6 May 2026 11:16:00 +0200 Subject: [PATCH 14/49] Switched to new GPU Codegen. --- pyproject.toml | 2 +- uv.lock | 11 ++++------- 2 files changed, 5 insertions(+), 8 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index c364c765d8..cbadfbf213 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -486,7 +486,7 @@ url = 'https://gridtools.github.io/pypi/' atlas4py = {index = "test.pypi"} dace = [ {git = "https://github.com/GridTools/dace", branch = "romanc/stree-v2", group = "dace-cartesian"}, - {index = "gridtools", group = "dace-next"} + {git = "https://github.com/spcl/dace", branch = "new-gpu-codegen-dev", group = "dace-next"} ] # -- versioningit -- diff --git a/uv.lock b/uv.lock index 19aaaa67a8..e358b459ba 100644 --- a/uv.lock +++ b/uv.lock @@ -1244,8 +1244,8 @@ dependencies = [ [[package]] name = "dace" -version = "43!2026.4.27" -source = { registry = "https://gridtools.github.io/pypi/" } +version = "1.0.0" +source = { git = "https://github.com/spcl/dace?branch=new-gpu-codegen-dev#d665388f79a9ce85750ea082ee438e76ce592b05" } resolution-markers = [ "python_full_version >= '3.14' and sys_platform == 'win32'", "python_full_version >= '3.14' and sys_platform == 'emscripten'", @@ -1276,9 +1276,6 @@ dependencies = [ { name = "sympy" }, { name = "typing-extensions" }, ] -wheels = [ - { url = "https://gridtools.github.io/pypi/dace/dace-43!2026.4.27-py3-none-any.whl", hash = "sha256:9098ceed412d287d575b2ed30cc90b754b966eda893698c55129a7bc5bd37d37" }, -] [[package]] name = "debugpy" @@ -1792,7 +1789,7 @@ dace-cartesian = [ { name = "dace", version = "1.0.0", source = { git = "https://github.com/GridTools/dace?branch=romanc%2Fstree-v2#d5fbadb626389e425fac5ed93d2a880811eca41f" } }, ] dace-next = [ - { name = "dace", version = "43!2026.4.27", source = { registry = "https://gridtools.github.io/pypi/" } }, + { name = "dace", version = "1.0.0", source = { git = "https://github.com/spcl/dace?branch=new-gpu-codegen-dev#d665388f79a9ce85750ea082ee438e76ce592b05" } }, ] dev = [ { name = "atlas4py" }, @@ -1962,7 +1959,7 @@ build = [ { name = "wheel", specifier = ">=0.33.6" }, ] dace-cartesian = [{ name = "dace", git = "https://github.com/GridTools/dace?branch=romanc%2Fstree-v2" }] -dace-next = [{ name = "dace", specifier = "==43!2026.4.27", index = "https://gridtools.github.io/pypi/", conflict = { package = "gt4py", group = "dace-next" } }] +dace-next = [{ name = "dace", git = "https://github.com/spcl/dace?branch=new-gpu-codegen-dev" }] dev = [ { name = "atlas4py", specifier = ">=0.41", index = "https://test.pypi.org/simple" }, { name = "coverage", extras = ["toml"], specifier = ">=7.6.1" }, From 574ed87cb85a25e7c1b936727338c978906b008a Mon Sep 17 00:00:00 2001 From: "Philip Mueller, CSCS" Date: Wed, 6 May 2026 11:42:05 +0200 Subject: [PATCH 15/49] Was this the error. --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index cbadfbf213..d1c7262353 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,7 @@ dace-cartesian = [ 'dace>=1.0.2' # refined in [tool.uv.sources] ] dace-next = [ - 'dace==43!2026.04.27' # uses custom index at 'https://github.com/GridTools/pypi' + 'dace==1.0.0' # uses custom index at 'https://github.com/GridTools/pypi' ] dev = [ {include-group = 'build'}, From 7a96e6b0bb06c5345f39b92382af05ba177f77e0 Mon Sep 17 00:00:00 2001 From: "Philip Mueller, CSCS" Date: Wed, 6 May 2026 11:48:55 +0200 Subject: [PATCH 16/49] This should be the thing. --- pyproject.toml | 4 ++-- uv.lock | 8 ++++---- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index d1c7262353..8bfc38aecb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,7 @@ dace-cartesian = [ 'dace>=1.0.2' # refined in [tool.uv.sources] ] dace-next = [ - 'dace==1.0.0' # uses custom index at 'https://github.com/GridTools/pypi' + 'dace==2.3.4' # uses custom index at 'https://github.com/GridTools/pypi' ] dev = [ {include-group = 'build'}, @@ -486,7 +486,7 @@ url = 'https://gridtools.github.io/pypi/' atlas4py = {index = "test.pypi"} dace = [ {git = "https://github.com/GridTools/dace", branch = "romanc/stree-v2", group = "dace-cartesian"}, - {git = "https://github.com/spcl/dace", branch = "new-gpu-codegen-dev", group = "dace-next"} + {git = "https://github.com/philip-paul-mueller/dace", branch = "phimuell__new-gpu-codegen-dev", group = "dace-next"} ] # -- versioningit -- diff --git a/uv.lock b/uv.lock index e358b459ba..1065704975 100644 --- a/uv.lock +++ b/uv.lock @@ -1244,8 +1244,8 @@ dependencies = [ [[package]] name = "dace" -version = "1.0.0" -source = { git = "https://github.com/spcl/dace?branch=new-gpu-codegen-dev#d665388f79a9ce85750ea082ee438e76ce592b05" } +version = "2.3.4" +source = { git = "https://github.com/philip-paul-mueller/dace?branch=phimuell__new-gpu-codegen-dev#877c1027a8fbbaf9e74879aae0719852bd4123dc" } resolution-markers = [ "python_full_version >= '3.14' and sys_platform == 'win32'", "python_full_version >= '3.14' and sys_platform == 'emscripten'", @@ -1789,7 +1789,7 @@ dace-cartesian = [ { name = "dace", version = "1.0.0", source = { git = "https://github.com/GridTools/dace?branch=romanc%2Fstree-v2#d5fbadb626389e425fac5ed93d2a880811eca41f" } }, ] dace-next = [ - { name = "dace", version = "1.0.0", source = { git = "https://github.com/spcl/dace?branch=new-gpu-codegen-dev#d665388f79a9ce85750ea082ee438e76ce592b05" } }, + { name = "dace", version = "2.3.4", source = { git = "https://github.com/philip-paul-mueller/dace?branch=phimuell__new-gpu-codegen-dev#877c1027a8fbbaf9e74879aae0719852bd4123dc" } }, ] dev = [ { name = "atlas4py" }, @@ -1959,7 +1959,7 @@ build = [ { name = "wheel", specifier = ">=0.33.6" }, ] dace-cartesian = [{ name = "dace", git = "https://github.com/GridTools/dace?branch=romanc%2Fstree-v2" }] -dace-next = [{ name = "dace", git = "https://github.com/spcl/dace?branch=new-gpu-codegen-dev" }] +dace-next = [{ name = "dace", git = "https://github.com/philip-paul-mueller/dace?branch=phimuell__new-gpu-codegen-dev" }] dev = [ { name = "atlas4py", specifier = ">=0.41", index = "https://test.pypi.org/simple" }, { name = "coverage", extras = ["toml"], specifier = ">=7.6.1" }, From fd463b3d657f70d99758f612274584f33a66bd1d Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 6 May 2026 15:44:35 +0200 Subject: [PATCH 17/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 82 +++++++++---------- .../dace_tests/test_dace_translation.py | 1 + 2 files changed, 42 insertions(+), 41 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 481c7b6138..fd0cc2f16a 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -280,10 +280,10 @@ def connect( dest_desc = self.result.dc_node.desc(self.state) write_edge = self.state.in_edges(self.result.dc_node)[0] - # Check the kind of node which writes the result - if isinstance(write_edge.src, dace_nodes.Tasklet): + if self.state.out_degree(self.result.dc_node) != 0: + remove_last_node = False + elif isinstance(write_edge.src, dace_nodes.Tasklet): # The temporary data written by a tasklet can be safely deleted. - assert map_exit is not None remove_last_node = True elif isinstance(write_edge.src, dace_nodes.NestedSDFG): if isinstance(dest_desc, dace.data.Scalar): @@ -726,9 +726,11 @@ def _visit_if_branch_arg( """ use_full_shape = False if isinstance(arg, (MemletExpr, ValueExpr)): + field_dims = [] arg_desc = arg.dc_node.desc(self.sdfg) arg_expr = arg elif isinstance(arg, IteratorExpr): + field_dims = [dim for dim, _ in arg.field_domain] arg_desc = arg.field.desc(self.sdfg) if deref_on_input_memlet: # If the iterator is just dereferenced inside the branch state, @@ -755,13 +757,21 @@ def _visit_if_branch_arg( inner_desc = dace.data.Scalar(arg_desc.dtype) else: # for list of values, we retrieve the local size from the corresponding offset - assert arg.gt_dtype.offset_type is not None - offset_provider_type = self.subgraph_builder.get_offset_provider_type( - arg.gt_dtype.offset_type.value + local_dim = arg.gt_dtype.offset_type + assert local_dim is not None + assert isinstance( + self.subgraph_builder.get_offset_provider_type(local_dim.value), + gtx_common.NeighborConnectivityType, ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + # find position of the local dimension in the field layout + assert isinstance(arg_desc, dace.data.Array) + assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in field_dims) + extended_dims = gtx_common.order_dimensions([*field_dims, local_dim]) + local_dim_pos = extended_dims.index(local_dim) inner_desc = dace.data.Array( - dtype=arg_desc.dtype, shape=[offset_provider_type.max_neighbors] + dtype=arg_desc.dtype, + shape=(arg_desc.shape[local_dim_pos],), + strides=(arg_desc.strides[local_dim_pos],), ) if param_name in if_sdfg.arrays: @@ -777,7 +787,7 @@ def _visit_if_branch_arg( else: return ValueExpr(inner_node, arg.gt_dtype) - def _visit_if_branch( + def _lower_if_state( self, if_sdfg: dace.SDFG, if_branch_state: dace.SDFGState, @@ -786,9 +796,10 @@ def _visit_if_branch( direct_deref_iterators: Iterable[str], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ - Helper method to visit an if-branch expression and lower it to a dataflow inside the given nested SDFG and state. + Helper method to visit an expression and lower it to a dataflow inside the given nested SDFG and state. - This function is called by `_visit_if()` for each if-branch. + This function is called by `_visit_if()` for the entry state (evaulation of + if-condition) and for each if-branch. Args: if_sdfg: The nested SDFG where the if expression is lowered. @@ -904,17 +915,6 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp assert len(node.args) == 3 - # evaluate the if-condition that will write to a boolean scalar node - condition_value = self.visit(node.args[0]) - assert ( - ( - isinstance(condition_value.gt_dtype, ts.ScalarType) - and condition_value.gt_dtype.kind == ts.ScalarKind.BOOL - ) - if isinstance(condition_value, (MemletExpr, ValueExpr)) - else (condition_value.dc_dtype == dace.dtypes.bool_) - ) - nsdfg = dace.SDFG(self.subgraph_builder.unique_nsdfg_name("if_stmt")) nsdfg.debuginfo = gtir_to_sdfg_utils.debug_info(node, default=self.sdfg.debuginfo) @@ -940,16 +940,6 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp # Use `None` for unconditional execution of else-branch, if the condition is not met. if_region.add_branch(None, else_body) - input_memlets: dict[str, MemletExpr | ValueExpr] = {} - nsdfg_symbols_mapping: Optional[dict[str, dace.symbol]] = None - - # define scalar or symbol for the condition value inside the nested SDFG - if isinstance(condition_value, SymbolExpr): - nsdfg.add_symbol("__cond", dace.dtypes.bool) - else: - nsdfg.add_scalar("__cond", dace.dtypes.bool) - input_memlets["__cond"] = condition_value - # Collect all field iterators that are shifted inside any of the then/else # branch expressions. Iterator shift expressions require the field argument # as iterator, therefore the corresponding array has to be passed with full @@ -959,7 +949,7 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp # be lowered outside the nested SDFG, so that just the local value (a scalar # or a list of values) is passed as input to the nested SDFG. shifted_iterator_symbols = set() - for branch_expr in node.args[1:3]: + for branch_expr in node.args: for shift_node in eve.walk_values(branch_expr).filter( lambda x: cpm.is_applied_shift(x) ): @@ -976,13 +966,28 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp if isinstance(sym_type, IteratorExpr) } direct_deref_iterators = ( - set(symbol_ref_utils.collect_symbol_refs(node.args[1:3], iterator_symbols)) + set(symbol_ref_utils.collect_symbol_refs(node.args, iterator_symbols)) - shifted_iterator_symbols ) + # collect all memlets that are needed as input to the nested SDFG + input_memlets: dict[str, MemletExpr | ValueExpr] = {} + + # evaluate the if-condition that will write to a boolean scalar node + in_edges, out_edge = self._lower_if_state( + nsdfg, entry_state, node.args[0], input_memlets, direct_deref_iterators + ) + for edge in in_edges: + edge.connect(map_entry=None) + assert isinstance(out_edge, DataflowOutputEdge) + condition_node = out_edge.result.dc_node + # write the boolean result to the '__cond' synbol on the interstate edge + nsdfg.add_symbol("__cond", dace.dtypes.bool) + nsdfg.out_edges(entry_state)[0].data.assignments["__cond"] = condition_node.data + for nstate, arg in zip([tstate, fstate], node.args[1:3]): # visit each if-branch in the corresponding state of the nested SDFG - in_edges, out_edges = self._visit_if_branch( + in_edges, out_edges = self._lower_if_state( nsdfg, nstate, arg, input_memlets, direct_deref_iterators ) for edge in in_edges: @@ -1012,16 +1017,11 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp if gtx_dace_args.is_connectivity_identifier(aname) and not adesc.transient } - # all free symbols are mapped to the symbols available in parent SDFG - nsdfg_symbols_mapping = {str(sym): sym for sym in nsdfg.free_symbols} - if isinstance(condition_value, SymbolExpr): - nsdfg_symbols_mapping["__cond"] = condition_value.value - nsdfg_node = self.state.add_nested_sdfg( nsdfg, inputs=(used_connectivities | input_memlets.keys()), outputs=outputs, - symbol_mapping=nsdfg_symbols_mapping, + symbol_mapping=None, # free symbols are mapped to symbols in the parent SDFG ) for inner, input_expr in input_memlets.items(): diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index 117aa1a92b..c3785c371e 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -65,6 +65,7 @@ def _translate_gtir_to_sdfg( # we use the SDFG hash in build cache to avoid clashes between CPU and GPU SDFGs return dace_wf_translation.DaCeTranslator( device_type=device_type, + apply_common_transform=False, auto_optimize=auto_optimize, auto_optimize_args=None, async_sdfg_call=async_sdfg_call, From 3b075ace72b5a2c1c79ff4bc12f315c86b3ddbe0 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 6 May 2026 15:58:14 +0200 Subject: [PATCH 18/49] edit --- src/gt4py/next/iterator/transforms/pass_manager.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index a3b440aed5..a25bceefde 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -138,7 +138,6 @@ def apply_common_transforms( unroll_reduce=False, common_subexpression_elimination=True, force_inline_lambda_args=False, - fuse_maps=True, transform_concat_where_to_as_fieldop=True, #: A dictionary mapping axes names to their length. See :func:`infer_domain.infer_expr` for #: more details. @@ -239,8 +238,7 @@ def apply_common_transforms( ir = NormalizeShifts().visit(ir) - if fuse_maps: - ir = FuseMaps(uids=uids).visit(ir) + ir = FuseMaps(uids=uids).visit(ir) ir = CollapseListGet().visit(ir) if unroll_reduce: From ae6db4e09bea081ec50b3e3559babc0585e6c614 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 6 May 2026 18:35:02 +0200 Subject: [PATCH 19/49] Fix stride of local dimension in lowering of if-expressions --- .../runners/dace/lowering/gtir_dataflow.py | 112 +++++++++++------- 1 file changed, 71 insertions(+), 41 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index da590d84e0..a2fe77cba9 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -681,9 +681,11 @@ def _visit_if_branch_arg( """ use_full_shape = False if isinstance(arg, (MemletExpr, ValueExpr)): + field_dims = [] arg_desc = arg.dc_node.desc(self.sdfg) arg_expr = arg elif isinstance(arg, IteratorExpr): + field_dims = [dim for dim, _ in arg.field_domain] arg_desc = arg.field.desc(self.sdfg) if deref_on_input_memlet: # If the iterator is just dereferenced inside the branch state, @@ -710,13 +712,21 @@ def _visit_if_branch_arg( inner_desc = dace.data.Scalar(arg_desc.dtype) else: # for list of values, we retrieve the local size from the corresponding offset - assert arg.gt_dtype.offset_type is not None - offset_provider_type = self.subgraph_builder.get_offset_provider_type( - arg.gt_dtype.offset_type.value + local_dim = arg.gt_dtype.offset_type + assert local_dim is not None + assert isinstance( + self.subgraph_builder.get_offset_provider_type(local_dim.value), + gtx_common.NeighborConnectivityType, ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + # find position of the local dimension in the field layout + assert isinstance(arg_desc, dace.data.Array) + assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in field_dims) + extended_dims = gtx_common.order_dimensions([*field_dims, local_dim]) + local_dim_pos = extended_dims.index(local_dim) inner_desc = dace.data.Array( - dtype=arg_desc.dtype, shape=[offset_provider_type.max_neighbors] + dtype=arg_desc.dtype, + shape=(arg_desc.shape[local_dim_pos],), + strides=(arg_desc.strides[local_dim_pos],), ) if param_name in if_sdfg.arrays: @@ -732,7 +742,7 @@ def _visit_if_branch_arg( else: return ValueExpr(inner_node, arg.gt_dtype) - def _visit_if_branch( + def _lower_if_state( self, if_sdfg: dace.SDFG, if_branch_state: dace.SDFGState, @@ -741,9 +751,10 @@ def _visit_if_branch( direct_deref_iterators: Iterable[str], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ - Helper method to visit an if-branch expression and lower it to a dataflow inside the given nested SDFG and state. + Helper method to visit an expression and lower it to a dataflow inside the given nested SDFG and state. - This function is called by `_visit_if()` for each if-branch. + This function is called by `_visit_if()` for the entry state (evaulation of + if-condition) and for each if-branch. Args: if_sdfg: The nested SDFG where the if expression is lowered. @@ -759,10 +770,15 @@ def _visit_if_branch( """ assert if_branch_state in if_sdfg.states() - lambda_args = [] - lambda_params = [] + lambda_args: list[MaybeNestedInTuple[IteratorExpr | DataExpr]] = [] + lambda_params: list[gtir.Sym] = [] for pname in symbol_ref_utils.collect_symbol_refs(expr, self.symbol_map.keys()): arg = self.symbol_map[pname] + if isinstance(arg, SymbolExpr): + psymbol = im.sym(pname, gtx_dace_args.as_itir_type(arg.dc_dtype)) + lambda_args.append(arg) + lambda_params.append(psymbol) + continue if isinstance(arg, tuple): ptype = get_tuple_type(arg) # type: ignore[arg-type] psymbol = im.sym(pname, ptype) @@ -781,7 +797,7 @@ def _visit_if_branch( ) )(psymbol_tree, arg) else: - psymbol = im.sym(pname, arg.gt_dtype) # type: ignore[union-attr] + psymbol = im.sym(pname, arg.gt_dtype) deref_on_input_memlet = pname in direct_deref_iterators inner_arg = self._visit_if_branch_arg( if_sdfg, @@ -854,20 +870,16 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp assert len(node.args) == 3 - # evaluate the if-condition that will write to a boolean scalar node - condition_value = self.visit(node.args[0]) - assert ( - ( - isinstance(condition_value.gt_dtype, ts.ScalarType) - and condition_value.gt_dtype.kind == ts.ScalarKind.BOOL - ) - if isinstance(condition_value, (MemletExpr, ValueExpr)) - else (condition_value.dc_dtype == dace.dtypes.bool_) - ) - nsdfg = dace.SDFG(self.subgraph_builder.unique_nsdfg_name("if_stmt")) nsdfg.debuginfo = gtir_to_sdfg_utils.debug_info(node, default=self.sdfg.debuginfo) + # add connectivities + for aname, adesc in self.sdfg.arrays.items(): + if gtx_dace_args.is_connectivity_identifier(aname): + adesc = adesc.clone() + adesc.transient = True + nsdfg.add_datadesc(aname, adesc) + # create states inside the nested SDFG for the if-branches if_region = dace.sdfg.state.ConditionalBlock("if") nsdfg.add_node(if_region, ensure_unique_name=True) @@ -883,16 +895,6 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp # Use `None` for unconditional execution of else-branch, if the condition is not met. if_region.add_branch(None, else_body) - input_memlets: dict[str, MemletExpr | ValueExpr] = {} - nsdfg_symbols_mapping: Optional[dict[str, dace.symbol]] = None - - # define scalar or symbol for the condition value inside the nested SDFG - if isinstance(condition_value, SymbolExpr): - nsdfg.add_symbol("__cond", dace.dtypes.bool) - else: - nsdfg.add_scalar("__cond", dace.dtypes.bool) - input_memlets["__cond"] = condition_value - # Collect all field iterators that are shifted inside any of the then/else # branch expressions. Iterator shift expressions require the field argument # as iterator, therefore the corresponding array has to be passed with full @@ -902,7 +904,7 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp # be lowered outside the nested SDFG, so that just the local value (a scalar # or a list of values) is passed as input to the nested SDFG. shifted_iterator_symbols = set() - for branch_expr in node.args[1:3]: + for branch_expr in node.args: for shift_node in eve.walk_values(branch_expr).filter( lambda x: cpm.is_applied_shift(x) ): @@ -919,13 +921,28 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp if isinstance(sym_type, IteratorExpr) } direct_deref_iterators = ( - set(symbol_ref_utils.collect_symbol_refs(node.args[1:3], iterator_symbols)) + set(symbol_ref_utils.collect_symbol_refs(node.args, iterator_symbols)) - shifted_iterator_symbols ) + # collect all memlets that are needed as input to the nested SDFG + input_memlets: dict[str, MemletExpr | ValueExpr] = {} + + # evaluate the if-condition that will write to a boolean scalar node + in_edges, out_edge = self._lower_if_state( + nsdfg, entry_state, node.args[0], input_memlets, direct_deref_iterators + ) + for edge in in_edges: + edge.connect(map_entry=None) + assert isinstance(out_edge, DataflowOutputEdge) + condition_node = out_edge.result.dc_node + # write the boolean result to the '__cond' synbol on the interstate edge + nsdfg.add_symbol("__cond", dace.dtypes.bool) + nsdfg.out_edges(entry_state)[0].data.assignments["__cond"] = condition_node.data + for nstate, arg in zip([tstate, fstate], node.args[1:3]): # visit each if-branch in the corresponding state of the nested SDFG - in_edges, out_edges = self._visit_if_branch( + in_edges, out_edges = self._lower_if_state( nsdfg, nstate, arg, input_memlets, direct_deref_iterators ) for edge in in_edges: @@ -948,15 +965,18 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp outputs = {outval.dc_node.data for outval in gtx_utils.flatten_nested_tuple((result,))} - # all free symbols are mapped to the symbols available in parent SDFG - nsdfg_symbols_mapping = {str(sym): sym for sym in nsdfg.free_symbols} - if isinstance(condition_value, SymbolExpr): - nsdfg_symbols_mapping["__cond"] = condition_value.value + # map the connectivities that were used inside the nested SDFG + used_connectivities = { + aname + for aname, adesc in nsdfg.arrays.items() + if gtx_dace_args.is_connectivity_identifier(aname) and not adesc.transient + } + nsdfg_node = self.state.add_nested_sdfg( nsdfg, - inputs=set(input_memlets.keys()), + inputs=(used_connectivities | input_memlets.keys()), outputs=outputs, - symbol_mapping=nsdfg_symbols_mapping, + symbol_mapping=None, # free symbols are mapped to symbols in the parent SDFG ) for inner, input_expr in input_memlets.items(): @@ -971,6 +991,16 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp self.sdfg.make_array_memlet(input_expr.dc_node.data), ) + for conn in used_connectivities: + desc = self.sdfg.data(conn) + desc.transient = False + self._add_input_data_edge( + self.state.add_access(conn), + dace_subsets.Range.from_array(desc), + nsdfg_node, + conn, + ) + return ( gtx_utils.tree_map(write_output_of_nested_sdfg_to_temporary)(result) if isinstance(result, tuple) From dbbfd526cf04d76c4342ba474215730c488a82a1 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 6 May 2026 18:35:17 +0200 Subject: [PATCH 20/49] edit --- .../program_processors/runners/dace/lowering/gtir_dataflow.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index a2fe77cba9..6690ca21ab 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -751,7 +751,8 @@ def _lower_if_state( direct_deref_iterators: Iterable[str], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ - Helper method to visit an expression and lower it to a dataflow inside the given nested SDFG and state. + Helper method to visit each argument of an if-expression and lower it to + a dataflow gragh inside the given nested SDFG and state. This function is called by `_visit_if()` for the entry state (evaulation of if-condition) and for each if-branch. From 2cd7728233ea3577ef2d67ad2c2f6f4ea04178c5 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 6 May 2026 18:42:30 +0200 Subject: [PATCH 21/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 6690ca21ab..c74c282df6 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -1921,6 +1921,9 @@ def visit_Literal(self, node: gtir.Literal) -> SymbolExpr: dc_dtype = gtx_dace_args.as_dace_type(node.type) return SymbolExpr(node.value, dc_dtype) + def visit_OffsetLiteral(self, node: gtir.OffsetLiteral) -> SymbolExpr: + return SymbolExpr(node.value, gtir_to_sdfg_types.INDEX_DTYPE) + def visit_SymRef(self, node: gtir.SymRef) -> MaybeNestedInTuple[IteratorExpr | DataExpr]: param = str(node.id) if param in self.symbol_map: @@ -1935,7 +1938,7 @@ def translate_lambda_to_dataflow( state: dace.SDFGState, sdfg_builder: gtir_to_sdfg.DataflowBuilder, node: gtir.Lambda, - args: Sequence[MaybeNestedInTuple[IteratorExpr | MemletExpr | ValueExpr]], + args: Sequence[MaybeNestedInTuple[IteratorExpr | DataExpr]], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ Entry point to visit a `Lambda` node and lower it to a dataflow graph, From 48a8f1ca8089462aed175e70dd72a84b05650086 Mon Sep 17 00:00:00 2001 From: "Philip Mueller, CSCS" Date: Thu, 7 May 2026 07:55:17 +0200 Subject: [PATCH 22/49] Let's try this fix. --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 8bfc38aecb..e084ae6092 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,7 @@ dace-cartesian = [ 'dace>=1.0.2' # refined in [tool.uv.sources] ] dace-next = [ - 'dace==2.3.4' # uses custom index at 'https://github.com/GridTools/pypi' + 'dace==2.3.5' # uses custom index at 'https://github.com/GridTools/pypi' ] dev = [ {include-group = 'build'}, From 6d52e245b9e6960bd50b9d5297ae9c4eb08cf1ea Mon Sep 17 00:00:00 2001 From: "Philip Mueller, CSCS" Date: Thu, 7 May 2026 07:56:44 +0200 Subject: [PATCH 23/49] Let's try this fix. --- uv.lock | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/uv.lock b/uv.lock index 1065704975..436decd62d 100644 --- a/uv.lock +++ b/uv.lock @@ -1244,8 +1244,8 @@ dependencies = [ [[package]] name = "dace" -version = "2.3.4" -source = { git = "https://github.com/philip-paul-mueller/dace?branch=phimuell__new-gpu-codegen-dev#877c1027a8fbbaf9e74879aae0719852bd4123dc" } +version = "2.3.5" +source = { git = "https://github.com/philip-paul-mueller/dace?branch=phimuell__new-gpu-codegen-dev#a62787d92d4ffe3f4586e8be1fdfc4169e791c17" } resolution-markers = [ "python_full_version >= '3.14' and sys_platform == 'win32'", "python_full_version >= '3.14' and sys_platform == 'emscripten'", @@ -1789,7 +1789,7 @@ dace-cartesian = [ { name = "dace", version = "1.0.0", source = { git = "https://github.com/GridTools/dace?branch=romanc%2Fstree-v2#d5fbadb626389e425fac5ed93d2a880811eca41f" } }, ] dace-next = [ - { name = "dace", version = "2.3.4", source = { git = "https://github.com/philip-paul-mueller/dace?branch=phimuell__new-gpu-codegen-dev#877c1027a8fbbaf9e74879aae0719852bd4123dc" } }, + { name = "dace", version = "2.3.5", source = { git = "https://github.com/philip-paul-mueller/dace?branch=phimuell__new-gpu-codegen-dev#a62787d92d4ffe3f4586e8be1fdfc4169e791c17" } }, ] dev = [ { name = "atlas4py" }, From ecd57805e7fbf3e89c2c0e73fb1cb3fb3ce23bbf Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 7 May 2026 09:33:11 +0200 Subject: [PATCH 24/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 71 ++++++++++--------- 1 file changed, 38 insertions(+), 33 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index c74c282df6..0ba5583e51 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -659,25 +659,25 @@ def _visit_deref(self, node: gtir.FunCall) -> DataExpr: field_desc.dtype, deref_node, connector_mapping["val"] ) - def _visit_if_branch_arg( + def _lower_if_state_arg( self, - if_sdfg: dace.SDFG, - if_branch_state: dace.SDFGState, + sdfg: dace.SDFG, + state: dace.SDFGState, param_name: str, arg: IteratorExpr | DataExpr, deref_on_input_memlet: bool, - if_sdfg_input_memlets: dict[str, MemletExpr | ValueExpr], + input_memlets: dict[str, MemletExpr | ValueExpr], ) -> IteratorExpr | ValueExpr: """ - Helper method to be called by `_visit_if_branch()` to visit the input arguments. + Helper method to be called by `_lower_if_state()` to visit the input arguments. Args: - if_sdfg: The nested SDFG where the if expression is lowered. - if_branch_state: The state inside the nested SDFG where the if branch is lowered. + sdfg: The SDFG where the if expression is lowered. + state: The state inside the SDFG where the argument is used. param_name: The parameter name of the input argument. arg: The input argument expression. deref_on_input_memlet: When True, the given iterator argument can be dereferenced on the input memlet. - if_sdfg_input_memlets: The memlets that provide input data to the nested SDFG, will be update inside this function. + input_memlets: The memlets that provide input data to the SDFG, which is updated inside this function. """ use_full_shape = False if isinstance(arg, (MemletExpr, ValueExpr)): @@ -709,7 +709,12 @@ def _visit_if_branch_arg( inner_desc = arg_desc.clone() inner_desc.transient = False elif isinstance(arg.gt_dtype, ts.ScalarType): - inner_desc = dace.data.Scalar(arg_desc.dtype) + if len(arg_desc.shape) == 1: + # TODO(edopao): cannot use a scalar because of an issue in gpu codegen, + # wich leads to compilation error: cannot convert 'const double' to 'const double*' + inner_desc = dace.data.Array(dtype=arg_desc.dtype, shape=(1,)) + else: + inner_desc = dace.data.Scalar(arg_desc.dtype) else: # for list of values, we retrieve the local size from the corresponding offset local_dim = arg.gt_dtype.offset_type @@ -729,14 +734,14 @@ def _visit_if_branch_arg( strides=(arg_desc.strides[local_dim_pos],), ) - if param_name in if_sdfg.arrays: + if param_name in sdfg.arrays: # the data desciptor was added by the visitor of the other branch expression - assert if_sdfg.data(param_name) == inner_desc + assert sdfg.data(param_name) == inner_desc else: - if_sdfg.add_datadesc(param_name, inner_desc) - if_sdfg_input_memlets[param_name] = arg_expr + sdfg.add_datadesc(param_name, inner_desc) + input_memlets[param_name] = arg_expr - inner_node = if_branch_state.add_access(param_name) + inner_node = state.add_access(param_name) if isinstance(arg, IteratorExpr) and use_full_shape: return IteratorExpr(inner_node, arg.gt_dtype, arg.field_domain, arg.indices) else: @@ -744,24 +749,24 @@ def _visit_if_branch_arg( def _lower_if_state( self, - if_sdfg: dace.SDFG, - if_branch_state: dace.SDFGState, + sdfg: dace.SDFG, + state: dace.SDFGState, expr: gtir.Expr, - if_sdfg_input_memlets: dict[str, MemletExpr | ValueExpr], + input_memlets: dict[str, MemletExpr | ValueExpr], direct_deref_iterators: Iterable[str], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ - Helper method to visit each argument of an if-expression and lower it to - a dataflow gragh inside the given nested SDFG and state. + Helper method to visit the subexpressions of an if-expression and lower them + to a dataflow graph inside the given nested SDFG and state. - This function is called by `_visit_if()` for the entry state (evaulation of + This function is called by `_visit_if()` for the entry state (evaluation of if-condition) and for each if-branch. Args: - if_sdfg: The nested SDFG where the if expression is lowered. - if_branch_state: The state inside the nested SDFG where the if branch is lowered. + sdfg: The SDFG where the if expression is lowered. + state: The state inside the SDFG where the if subexpression is lowered. expr: The if branch expression to lower. - if_sdfg_input_memlets: The memlets that provide input data to the nested SDFG, will be update inside this function. + input_memlets: The memlets that provide input data to the SDFG, which are update inside this function. direct_deref_iterators: Fields that are accessed with direct iterator deref, without any shift. Returns: @@ -769,7 +774,7 @@ def _lower_if_state( - the list of input edges for the parent dataflow - the tree representation of output data, in the form of a tuple of data edges. """ - assert if_branch_state in if_sdfg.states() + assert state in sdfg.states() lambda_args: list[MaybeNestedInTuple[IteratorExpr | DataExpr]] = [] lambda_params: list[gtir.Sym] = [] @@ -787,26 +792,26 @@ def _lower_if_state( deref_on_input_memlet = pname in direct_deref_iterators inner_arg = gtx_utils.tree_map( lambda tsym, targ, deref_on_input_memlet=deref_on_input_memlet: ( - self._visit_if_branch_arg( - if_sdfg, - if_branch_state, + self._lower_if_state_arg( + sdfg, + state, str(tsym.id), targ, deref_on_input_memlet, - if_sdfg_input_memlets, + input_memlets, ) ) )(psymbol_tree, arg) else: psymbol = im.sym(pname, arg.gt_dtype) deref_on_input_memlet = pname in direct_deref_iterators - inner_arg = self._visit_if_branch_arg( - if_sdfg, - if_branch_state, + inner_arg = self._lower_if_state_arg( + sdfg, + state, pname, arg, deref_on_input_memlet, - if_sdfg_input_memlets, + input_memlets, ) lambda_args.append(inner_arg) lambda_params.append(psymbol) @@ -814,7 +819,7 @@ def _lower_if_state( # visit each branch of the if-statement as if it was a Lambda node lambda_node = gtir.Lambda(params=lambda_params, expr=expr) input_edges, output_tree = translate_lambda_to_dataflow( - if_sdfg, if_branch_state, self.subgraph_builder, lambda_node, lambda_args + sdfg, state, self.subgraph_builder, lambda_node, lambda_args ) return input_edges, output_tree From a8fa800672ef5095702e7463ec814254f76cef2e Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 7 May 2026 12:04:52 +0200 Subject: [PATCH 25/49] edit --- .../program_processors/runners/dace/lowering/gtir_dataflow.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 0ba5583e51..904c097935 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -709,7 +709,7 @@ def _lower_if_state_arg( inner_desc = arg_desc.clone() inner_desc.transient = False elif isinstance(arg.gt_dtype, ts.ScalarType): - if len(arg_desc.shape) == 1: + if isinstance(arg, MemletExpr) and len(arg_expr.gt_field.dims) == 1: # TODO(edopao): cannot use a scalar because of an issue in gpu codegen, # wich leads to compilation error: cannot convert 'const double' to 'const double*' inner_desc = dace.data.Array(dtype=arg_desc.dtype, shape=(1,)) From 80c45e6432589964475d184f570be7604f087c15 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 7 May 2026 12:11:15 +0200 Subject: [PATCH 26/49] edit --- .../program_processors/runners/dace/lowering/gtir_dataflow.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 904c097935..b7eb66f849 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -709,7 +709,7 @@ def _lower_if_state_arg( inner_desc = arg_desc.clone() inner_desc.transient = False elif isinstance(arg.gt_dtype, ts.ScalarType): - if isinstance(arg, MemletExpr) and len(arg_expr.gt_field.dims) == 1: + if isinstance(arg, MemletExpr) and len(arg.gt_field.dims) == 1: # TODO(edopao): cannot use a scalar because of an issue in gpu codegen, # wich leads to compilation error: cannot convert 'const double' to 'const double*' inner_desc = dace.data.Array(dtype=arg_desc.dtype, shape=(1,)) From ba91d43b05f45684afd85128277ce695f6bf3315 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 7 May 2026 13:34:00 +0200 Subject: [PATCH 27/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 26 +++++++++---------- 1 file changed, 13 insertions(+), 13 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index b7eb66f849..7373b5a560 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -673,11 +673,11 @@ def _lower_if_state_arg( Args: sdfg: The SDFG where the if expression is lowered. - state: The state inside the SDFG where the argument is used. - param_name: The parameter name of the input argument. - arg: The input argument expression. + state: The state inside the given SDFG where the argument is used. + param_name: The name of the input argument. + arg: The expression corresponding to the input argument. deref_on_input_memlet: When True, the given iterator argument can be dereferenced on the input memlet. - input_memlets: The memlets that provide input data to the SDFG, which is updated inside this function. + input_memlets: The memlets that provide input data to the SDFG, will be updated inside this function. """ use_full_shape = False if isinstance(arg, (MemletExpr, ValueExpr)): @@ -710,8 +710,8 @@ def _lower_if_state_arg( inner_desc.transient = False elif isinstance(arg.gt_dtype, ts.ScalarType): if isinstance(arg, MemletExpr) and len(arg.gt_field.dims) == 1: - # TODO(edopao): cannot use a scalar because of an issue in gpu codegen, - # wich leads to compilation error: cannot convert 'const double' to 'const double*' + # TODO(edopao): we cannot use a scalar because of an issue in gpu codegen, + # which leads to compilation error: cannot convert 'const double' to 'const double*' inner_desc = dace.data.Array(dtype=arg_desc.dtype, shape=(1,)) else: inner_desc = dace.data.Scalar(arg_desc.dtype) @@ -756,17 +756,17 @@ def _lower_if_state( direct_deref_iterators: Iterable[str], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ - Helper method to visit the subexpressions of an if-expression and lower them - to a dataflow graph inside the given nested SDFG and state. + Helper method to visit a subexpression inside an if-node and lower it + to a dataflow graph inside the given SDFG and state. This function is called by `_visit_if()` for the entry state (evaluation of - if-condition) and for each if-branch. + if-condition) and for each if-branch state. Args: sdfg: The SDFG where the if expression is lowered. - state: The state inside the SDFG where the if subexpression is lowered. - expr: The if branch expression to lower. - input_memlets: The memlets that provide input data to the SDFG, which are update inside this function. + state: The state inside the given SDFG where the subexpression is lowered. + expr: The subexpression to lower. + input_memlets: The memlets that provide input data to the SDFG, will be updated inside this function. direct_deref_iterators: Fields that are accessed with direct iterator deref, without any shift. Returns: @@ -942,7 +942,7 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp edge.connect(map_entry=None) assert isinstance(out_edge, DataflowOutputEdge) condition_node = out_edge.result.dc_node - # write the boolean result to the '__cond' synbol on the interstate edge + # write the boolean result to the '__cond' symbol on the interstate edge nsdfg.add_symbol("__cond", dace.dtypes.bool) nsdfg.out_edges(entry_state)[0].data.assignments["__cond"] = condition_node.data From f56010746631697427e6c12edd211f04436dd969 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 7 May 2026 14:29:02 +0200 Subject: [PATCH 28/49] apply review comments --- .../runners/dace/lowering/gtir_dataflow.py | 85 ++++++++++--------- 1 file changed, 46 insertions(+), 39 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 7373b5a560..0e2d2142e3 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -659,7 +659,7 @@ def _visit_deref(self, node: gtir.FunCall) -> DataExpr: field_desc.dtype, deref_node, connector_mapping["val"] ) - def _lower_if_state_arg( + def _visit_if_branch_arg( self, sdfg: dace.SDFG, state: dace.SDFGState, @@ -669,13 +669,13 @@ def _lower_if_state_arg( input_memlets: dict[str, MemletExpr | ValueExpr], ) -> IteratorExpr | ValueExpr: """ - Helper method to be called by `_lower_if_state()` to visit the input arguments. + Helper method to be called by `_visit_if_branch()` to visit the input arguments. Args: sdfg: The SDFG where the if expression is lowered. - state: The state inside the given SDFG where the argument is used. - param_name: The name of the input argument. - arg: The expression corresponding to the input argument. + state: The state inside the given SDFG where the if branch is lowered. + param_name: The parameter name of the input argument. + arg: The input argument expression. deref_on_input_memlet: When True, the given iterator argument can be dereferenced on the input memlet. input_memlets: The memlets that provide input data to the SDFG, will be updated inside this function. """ @@ -747,7 +747,7 @@ def _lower_if_state_arg( else: return ValueExpr(inner_node, arg.gt_dtype) - def _lower_if_state( + def _visit_if_branch( self, sdfg: dace.SDFG, state: dace.SDFGState, @@ -756,16 +756,14 @@ def _lower_if_state( direct_deref_iterators: Iterable[str], ) -> tuple[list[DataflowInputEdge], MaybeNestedInTuple[DataflowOutputEdge]]: """ - Helper method to visit a subexpression inside an if-node and lower it - to a dataflow graph inside the given SDFG and state. + Helper method to visit an if-branch expression and lower it to a dataflow inside the given nested SDFG and state. - This function is called by `_visit_if()` for the entry state (evaluation of - if-condition) and for each if-branch state. + This function is called by `_visit_if()` for each if-branch. Args: sdfg: The SDFG where the if expression is lowered. - state: The state inside the given SDFG where the subexpression is lowered. - expr: The subexpression to lower. + state: The state inside the given SDFG where the if branch is lowered. + expr: The if branch expression to lower. input_memlets: The memlets that provide input data to the SDFG, will be updated inside this function. direct_deref_iterators: Fields that are accessed with direct iterator deref, without any shift. @@ -780,19 +778,18 @@ def _lower_if_state( lambda_params: list[gtir.Sym] = [] for pname in symbol_ref_utils.collect_symbol_refs(expr, self.symbol_map.keys()): arg = self.symbol_map[pname] + inner_arg: MaybeNestedInTuple[IteratorExpr | DataExpr] if isinstance(arg, SymbolExpr): psymbol = im.sym(pname, gtx_dace_args.as_itir_type(arg.dc_dtype)) - lambda_args.append(arg) - lambda_params.append(psymbol) - continue - if isinstance(arg, tuple): + inner_arg = arg + elif isinstance(arg, tuple): ptype = get_tuple_type(arg) # type: ignore[arg-type] psymbol = im.sym(pname, ptype) psymbol_tree = gtir_to_sdfg_utils.make_symbol_tree(pname, ptype) deref_on_input_memlet = pname in direct_deref_iterators inner_arg = gtx_utils.tree_map( lambda tsym, targ, deref_on_input_memlet=deref_on_input_memlet: ( - self._lower_if_state_arg( + self._visit_if_branch_arg( sdfg, state, str(tsym.id), @@ -805,7 +802,7 @@ def _lower_if_state( else: psymbol = im.sym(pname, arg.gt_dtype) deref_on_input_memlet = pname in direct_deref_iterators - inner_arg = self._lower_if_state_arg( + inner_arg = self._visit_if_branch_arg( sdfg, state, pname, @@ -876,6 +873,17 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp assert len(node.args) == 3 + # evaluate the if-condition that will write to a boolean scalar node + condition_value = self.visit(node.args[0]) + assert ( + ( + isinstance(condition_value.gt_dtype, ts.ScalarType) + and condition_value.gt_dtype.kind == ts.ScalarKind.BOOL + ) + if isinstance(condition_value, (MemletExpr, ValueExpr)) + else (condition_value.dc_dtype == dace.dtypes.bool_) + ) + nsdfg = dace.SDFG(self.subgraph_builder.unique_nsdfg_name("if_stmt")) nsdfg.debuginfo = gtir_to_sdfg_utils.debug_info(node, default=self.sdfg.debuginfo) @@ -901,6 +909,16 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp # Use `None` for unconditional execution of else-branch, if the condition is not met. if_region.add_branch(None, else_body) + input_memlets: dict[str, MemletExpr | ValueExpr] = {} + nsdfg_symbols_mapping: Optional[dict[str, dace.symbol]] = None + + # define scalar or symbol for the condition value inside the nested SDFG + if isinstance(condition_value, SymbolExpr): + nsdfg.add_symbol("__cond", dace.dtypes.bool) + else: + nsdfg.add_scalar("__cond", dace.dtypes.bool) + input_memlets["__cond"] = condition_value + # Collect all field iterators that are shifted inside any of the then/else # branch expressions. Iterator shift expressions require the field argument # as iterator, therefore the corresponding array has to be passed with full @@ -910,7 +928,7 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp # be lowered outside the nested SDFG, so that just the local value (a scalar # or a list of values) is passed as input to the nested SDFG. shifted_iterator_symbols = set() - for branch_expr in node.args: + for branch_expr in node.args[1:3]: for shift_node in eve.walk_values(branch_expr).filter( lambda x: cpm.is_applied_shift(x) ): @@ -927,28 +945,13 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp if isinstance(sym_type, IteratorExpr) } direct_deref_iterators = ( - set(symbol_ref_utils.collect_symbol_refs(node.args, iterator_symbols)) + set(symbol_ref_utils.collect_symbol_refs(node.args[1:3], iterator_symbols)) - shifted_iterator_symbols ) - # collect all memlets that are needed as input to the nested SDFG - input_memlets: dict[str, MemletExpr | ValueExpr] = {} - - # evaluate the if-condition that will write to a boolean scalar node - in_edges, out_edge = self._lower_if_state( - nsdfg, entry_state, node.args[0], input_memlets, direct_deref_iterators - ) - for edge in in_edges: - edge.connect(map_entry=None) - assert isinstance(out_edge, DataflowOutputEdge) - condition_node = out_edge.result.dc_node - # write the boolean result to the '__cond' symbol on the interstate edge - nsdfg.add_symbol("__cond", dace.dtypes.bool) - nsdfg.out_edges(entry_state)[0].data.assignments["__cond"] = condition_node.data - for nstate, arg in zip([tstate, fstate], node.args[1:3]): # visit each if-branch in the corresponding state of the nested SDFG - in_edges, out_edges = self._lower_if_state( + in_edges, out_edges = self._visit_if_branch( nsdfg, nstate, arg, input_memlets, direct_deref_iterators ) for edge in in_edges: @@ -978,11 +981,15 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp if gtx_dace_args.is_connectivity_identifier(aname) and not adesc.transient } + # all free symbols are mapped to the symbols available in parent SDFG + nsdfg_symbols_mapping = {str(sym): sym for sym in nsdfg.free_symbols} + if isinstance(condition_value, SymbolExpr): + nsdfg_symbols_mapping["__cond"] = condition_value.value nsdfg_node = self.state.add_nested_sdfg( nsdfg, - inputs=(used_connectivities | input_memlets.keys()), - outputs=outputs, - symbol_mapping=None, # free symbols are mapped to symbols in the parent SDFG + inputs=sorted(used_connectivities | input_memlets.keys()), + outputs=sorted(outputs), + symbol_mapping=nsdfg_symbols_mapping, ) for inner, input_expr in input_memlets.items(): From c174e5ac10c50fef89445228a3a19cc4e59b280d Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Fri, 8 May 2026 11:15:01 +0200 Subject: [PATCH 29/49] fix --- .../runners/dace/lowering/gtir_dataflow.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 0e2d2142e3..fe54255918 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -837,9 +837,16 @@ def _visit_if_branch_result( # If the result is currently written to a transient node, inside the nested SDFG, # we need to allocate a non-transient data node. result_desc = edge.result.dc_node.desc(sdfg) - output_desc = result_desc.clone() - output_desc.transient = False - output_data = sdfg.add_datadesc(output_data, output_desc, find_new_name=True) + if isinstance(sym.type, ts.ScalarType) and isinstance(result_desc, dace.data.Array): + # TODO(edopao): a scalar should not be represented as an array, but + # currently this can happen because of an issue workaround, see the + # todo comment above in `_visit_if_branch_arg()`. + assert len(result_desc.shape) == 1 and result_desc.shape[0] == 1 + _, output_desc = sdfg.add_scalar(output_data, result_desc.dtype) + else: + output_desc = result_desc.clone() + output_desc.transient = False + output_data = sdfg.add_datadesc(output_data, output_desc, find_new_name=True) output_node = state.add_access(output_data) state.add_nedge( edge.result.dc_node, From 07a587700c1b0aaa66b2b6c46c139aafab17ea9c Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Fri, 8 May 2026 12:14:54 +0200 Subject: [PATCH 30/49] remove workaround --- .../runners/dace/lowering/gtir_dataflow.py | 20 ++++--------------- 1 file changed, 4 insertions(+), 16 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 221181ad2f..52a1f2ff20 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -754,12 +754,7 @@ def _visit_if_branch_arg( inner_desc = arg_desc.clone() inner_desc.transient = False elif isinstance(arg.gt_dtype, ts.ScalarType): - if isinstance(arg, MemletExpr) and len(arg.gt_field.dims) == 1: - # TODO(edopao): we cannot use a scalar because of an issue in gpu codegen, - # which leads to compilation error: cannot convert 'const double' to 'const double*' - inner_desc = dace.data.Array(dtype=arg_desc.dtype, shape=(1,)) - else: - inner_desc = dace.data.Scalar(arg_desc.dtype) + inner_desc = dace.data.Scalar(arg_desc.dtype) else: # for list of values, we retrieve the local size from the corresponding offset local_dim = arg.gt_dtype.offset_type @@ -882,16 +877,9 @@ def _visit_if_branch_result( # If the result is currently written to a transient node, inside the nested SDFG, # we need to allocate a non-transient data node. result_desc = edge.result.dc_node.desc(sdfg) - if isinstance(sym.type, ts.ScalarType) and isinstance(result_desc, dace.data.Array): - # TODO(edopao): a scalar should not be represented as an array, but - # currently this can happen because of an issue workaround, see the - # todo comment above in `_visit_if_branch_arg()`. - assert len(result_desc.shape) == 1 and result_desc.shape[0] == 1 - _, output_desc = sdfg.add_scalar(output_data, result_desc.dtype) - else: - output_desc = result_desc.clone() - output_desc.transient = False - output_data = sdfg.add_datadesc(output_data, output_desc, find_new_name=True) + output_desc = result_desc.clone() + output_desc.transient = False + output_data = sdfg.add_datadesc(output_data, output_desc, find_new_name=True) output_node = state.add_access(output_data) state.add_nedge( edge.result.dc_node, From 705f512652107ec1329edbce0a00eeec0ca58f73 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Tue, 12 May 2026 10:33:11 +0200 Subject: [PATCH 31/49] re-enable test --- .../iterator_tests/test_vertical_advection.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py index e98e820f14..ffae942a15 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py @@ -99,9 +99,6 @@ def fen_solve_tridiag2(i_size, j_size, k_size, a, b, c, d, x): def test_tridiag(fencil, tridiag_reference, program_processor): program_processor, validate = program_processor - if isinstance(program_processor, backend.Backend) and "dace" in program_processor.name: - pytest.xfail("Dace ITIR backend doesn't support the IR format used in this test.") - a, b, c, d, x = tridiag_reference shape = a.shape as_3d_field = gtx.as_field.partial([IDim, JDim, KDim]) From 923ab93dd4aa815e581db28e4326b6478bcf82d9 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Tue, 12 May 2026 13:37:01 +0200 Subject: [PATCH 32/49] preserve order of connectivities --- .../program_processors/runners/dace/lowering/gtir_dataflow.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 52a1f2ff20..e524aa6091 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -1015,11 +1015,11 @@ def write_output_of_nested_sdfg_to_temporary(inner_value: ValueExpr) -> ValueExp outputs = {outval.dc_node.data for outval in gtx_utils.flatten_nested_tuple((result,))} # map the connectivities that were used inside the nested SDFG - used_connectivities = { + used_connectivities = [ aname for aname, adesc in nsdfg.arrays.items() if gtx_dace_args.is_connectivity_identifier(aname) and not adesc.transient - } + ] # all free symbols are mapped to the symbols available in parent SDFG nsdfg_symbols_mapping = {str(sym): sym for sym in nsdfg.free_symbols} From af56259e6355ba591afc47aea1349fe7f140560b Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Tue, 12 May 2026 15:33:28 +0200 Subject: [PATCH 33/49] remove workaround --- .../runners/dace/lowering/gtir_dataflow.py | 20 ++++--------------- 1 file changed, 4 insertions(+), 16 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 05cbe4a798..e524aa6091 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -754,12 +754,7 @@ def _visit_if_branch_arg( inner_desc = arg_desc.clone() inner_desc.transient = False elif isinstance(arg.gt_dtype, ts.ScalarType): - if isinstance(arg, MemletExpr) and len(arg.gt_field.dims) == 1: - # TODO(edopao): we cannot use a scalar because of an issue in gpu codegen, - # which leads to compilation error: cannot convert 'const double' to 'const double*' - inner_desc = dace.data.Array(dtype=arg_desc.dtype, shape=(1,)) - else: - inner_desc = dace.data.Scalar(arg_desc.dtype) + inner_desc = dace.data.Scalar(arg_desc.dtype) else: # for list of values, we retrieve the local size from the corresponding offset local_dim = arg.gt_dtype.offset_type @@ -882,16 +877,9 @@ def _visit_if_branch_result( # If the result is currently written to a transient node, inside the nested SDFG, # we need to allocate a non-transient data node. result_desc = edge.result.dc_node.desc(sdfg) - if isinstance(sym.type, ts.ScalarType) and isinstance(result_desc, dace.data.Array): - # TODO(edopao): a scalar should not be represented as an array, but - # currently this can happen because of an issue workaround, see the - # todo comment above in `_visit_if_branch_arg()`. - assert len(result_desc.shape) == 1 and result_desc.shape[0] == 1 - _, output_desc = sdfg.add_scalar(output_data, result_desc.dtype) - else: - output_desc = result_desc.clone() - output_desc.transient = False - output_data = sdfg.add_datadesc(output_data, output_desc, find_new_name=True) + output_desc = result_desc.clone() + output_desc.transient = False + output_data = sdfg.add_datadesc(output_data, output_desc, find_new_name=True) output_node = state.add_access(output_data) state.add_nedge( edge.result.dc_node, From 39fca693cefb177f23f8027e638ba4a25a05e3ab Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Tue, 12 May 2026 15:53:07 +0200 Subject: [PATCH 34/49] Revert "re-enable test" This reverts commit 705f512652107ec1329edbce0a00eeec0ca58f73. --- .../iterator_tests/test_vertical_advection.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py index ffae942a15..e98e820f14 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_vertical_advection.py @@ -99,6 +99,9 @@ def fen_solve_tridiag2(i_size, j_size, k_size, a, b, c, d, x): def test_tridiag(fencil, tridiag_reference, program_processor): program_processor, validate = program_processor + if isinstance(program_processor, backend.Backend) and "dace" in program_processor.name: + pytest.xfail("Dace ITIR backend doesn't support the IR format used in this test.") + a, b, c, d, x = tridiag_reference shape = a.shape as_3d_field = gtx.as_field.partial([IDim, JDim, KDim]) From 1621acfb7ca30a0858e79642d236ba1216ca6cda Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 13 May 2026 16:03:00 +0200 Subject: [PATCH 35/49] fix stride gpu transformation --- .../runners/dace/transformations/strides.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py index 0840f77755..17c92e739c 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py @@ -6,6 +6,8 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import functools +import operator from typing import Optional, TypeAlias import dace @@ -497,6 +499,10 @@ def _gt_map_strides_into_nested_sdfg( # A scalar does not have a stride that must be propagated. return + if functools.reduce(operator.mul, inner_shape) == 1: + # Arrays with size 1 in all dimensions do not need strides. + return + # Now determine the new stride that is needed on the inside. new_strides: list = [] if len(outer_shape) == len(inner_shape): From 1217498c387f722f2018c5a06715655e057bb799 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 13 May 2026 16:30:09 +0200 Subject: [PATCH 36/49] edit --- .../runners/dace/transformations/strides.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py index 17c92e739c..9993bf3cbd 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py @@ -499,8 +499,8 @@ def _gt_map_strides_into_nested_sdfg( # A scalar does not have a stride that must be propagated. return - if functools.reduce(operator.mul, inner_shape) == 1: - # Arrays with size 1 in all dimensions do not need strides. + if (functools.reduce(operator.mul, inner_shape) == 1) == True: # noqa: E712 [true-false-comparison] # SymPy comparison + # Arrays with size 1 in all dimensions don't use strides. return # Now determine the new stride that is needed on the inside. From f305c1c46d9e17fb4baa01d21926196cde93d2a6 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 13 May 2026 16:52:32 +0200 Subject: [PATCH 37/49] edit --- .../program_processors/runners/dace/transformations/strides.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py index 9993bf3cbd..f1cb32aa14 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py @@ -499,7 +499,7 @@ def _gt_map_strides_into_nested_sdfg( # A scalar does not have a stride that must be propagated. return - if (functools.reduce(operator.mul, inner_shape) == 1) == True: # noqa: E712 [true-false-comparison] # SymPy comparison + if (functools.reduce(operator.mul, inner_shape, 1) == 1) == True: # noqa: E712 [true-false-comparison] # SymPy comparison # Arrays with size 1 in all dimensions don't use strides. return From fb7c6f77674a133801920908536ece7a8d1c293c Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 13 May 2026 17:32:59 +0200 Subject: [PATCH 38/49] edit --- .../runners/dace/transformations/strides.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py index f1cb32aa14..03b13b3d6b 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py @@ -6,8 +6,7 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause -import functools -import operator +import warnings from typing import Optional, TypeAlias import dace @@ -499,10 +498,6 @@ def _gt_map_strides_into_nested_sdfg( # A scalar does not have a stride that must be propagated. return - if (functools.reduce(operator.mul, inner_shape, 1) == 1) == True: # noqa: E712 [true-false-comparison] # SymPy comparison - # Arrays with size 1 in all dimensions don't use strides. - return - # Now determine the new stride that is needed on the inside. new_strides: list = [] if len(outer_shape) == len(inner_shape): @@ -533,7 +528,9 @@ def _gt_map_strides_into_nested_sdfg( raise NotImplementedError("NestedSDFGs can not be used to increase the rank.") if len(new_strides) != len(inner_shape): - raise ValueError("Failed to compute the inner strides.") + # It might still be possible to access an array at index 0. + warnings.warn("Failed to compute the inner strides.", stacklevel=2) + return # For the strides of the arrays inside the nested SDFG we will create a new unique # symbol which is initialized, through the symbol mapping, to the value of this From e124ebfa2ae0537cf1aade2617d1f0f84394ffad Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 13 May 2026 17:48:16 +0200 Subject: [PATCH 39/49] edit --- .../runners/dace/transformations/strides.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py index 03b13b3d6b..168f847ee8 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/strides.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py @@ -528,7 +528,9 @@ def _gt_map_strides_into_nested_sdfg( raise NotImplementedError("NestedSDFGs can not be used to increase the rank.") if len(new_strides) != len(inner_shape): - # It might still be possible to access an array at index 0. + # It could still be possible to access an array at index 0. Consider a memlet + # which only writes index 0 to the inner shape (dim_oinflow == 1), although + # the inner shape is larger than 1, but we only read index 0 inside the SDFG. warnings.warn("Failed to compute the inner strides.", stacklevel=2) return From c328157a9ff0ce50fdece680d5b6bd32add550b4 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 13 May 2026 18:57:00 +0200 Subject: [PATCH 40/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index e524aa6091..59da7016a4 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -277,6 +277,10 @@ def connect( outside data container is removed, the caller is responsible to propagate the strides of the destination array to the array inside the nested SDFG. """ + # We doo not allow removing the last node of the dataflow inside a map scope, + # because the same reults could be written to multiple destiation nodes ouside + # the map scope, in case of field operators returning tuples. + allow_removal_of_last_node: Final[bool] = False dest_desc = self.result.dc_node.desc(self.state) write_edge = self.state.in_edges(self.result.dc_node)[0] @@ -301,7 +305,7 @@ def connect( else: remove_last_node = False - if remove_last_node: + if allow_removal_of_last_node and remove_last_node: src_node = write_edge.src src_node_connector = write_edge.src_conn src_subset = write_edge.data.src_subset From e60bd289b57c5667c1aef8f3baf72574739b2569 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Wed, 13 May 2026 19:03:59 +0200 Subject: [PATCH 41/49] edit --- .../program_processors/runners/dace/lowering/gtir_dataflow.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 59da7016a4..3bac7bf57a 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -277,7 +277,7 @@ def connect( outside data container is removed, the caller is responsible to propagate the strides of the destination array to the array inside the nested SDFG. """ - # We doo not allow removing the last node of the dataflow inside a map scope, + # We don't allow removing the last node of the dataflow inside a map scope, # because the same reults could be written to multiple destiation nodes ouside # the map scope, in case of field operators returning tuples. allow_removal_of_last_node: Final[bool] = False From 609aa97eaedfe2e4a0fa34d349ac8221a6b6355c Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Fri, 15 May 2026 10:13:57 +0200 Subject: [PATCH 42/49] edit --- .../runners/dace/lowering/gtir_dataflow.py | 42 +++++++++---------- .../dace/lowering/gtir_to_sdfg_primitives.py | 20 +++++++-- .../dace/lowering/gtir_to_sdfg_scan.py | 18 +++++++- 3 files changed, 54 insertions(+), 26 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py index 3bac7bf57a..c26e8625ec 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_dataflow.py @@ -269,6 +269,7 @@ def connect( map_exit: Optional[dace_nodes.MapExit], dest: dace_nodes.AccessNode, dest_subset: dace_subsets.Range, + allow_removal_of_last_node: bool, ) -> bool: """Create a connection to the `dest` node, writing the given `dest_subset`. @@ -277,35 +278,34 @@ def connect( outside data container is removed, the caller is responsible to propagate the strides of the destination array to the array inside the nested SDFG. """ - # We don't allow removing the last node of the dataflow inside a map scope, - # because the same reults could be written to multiple destiation nodes ouside - # the map scope, in case of field operators returning tuples. - allow_removal_of_last_node: Final[bool] = False dest_desc = self.result.dc_node.desc(self.state) write_edge = self.state.in_edges(self.result.dc_node)[0] - if self.state.out_degree(self.result.dc_node) != 0: - remove_last_node = False - elif isinstance(write_edge.src, dace_nodes.Tasklet): - # The temporary data written by a tasklet can be safely deleted. - remove_last_node = True - elif isinstance(write_edge.src, dace_nodes.NestedSDFG): - if isinstance(dest_desc, dace.data.Scalar): - # We keep scalar temporary storage, as a general rule, since it - # does not affect performance of the generated code. This scalar - # is only required in some cases, e.g. for nested SDFGs implementing - # reduction, which use a WCR memlet for the reduction operation. + if allow_removal_of_last_node: + if self.state.out_degree(self.result.dc_node) != 0: remove_last_node = False - else: - # We remove the transient array on the output connection of a nested - # SDFG and write directly to the destination node. - # The caller is responsible to propagate the strides of the destination - # array to the array inside the nested SDFG. + elif isinstance(write_edge.src, dace_nodes.Tasklet): + # The temporary data written by a tasklet can be safely deleted. remove_last_node = True + elif isinstance(write_edge.src, dace_nodes.NestedSDFG): + if isinstance(dest_desc, dace.data.Scalar): + # We keep scalar temporary storage, as a general rule, since it + # does not affect performance of the generated code. This scalar + # is only required in some cases, e.g. for nested SDFGs implementing + # reduction, which use a WCR memlet for the reduction operation. + remove_last_node = False + else: + # We remove the transient array on the output connection of a nested + # SDFG and write directly to the destination node. + # The caller is responsible to propagate the strides of the destination + # array to the array inside the nested SDFG. + remove_last_node = True + else: + remove_last_node = False else: remove_last_node = False - if allow_removal_of_last_node and remove_last_node: + if remove_last_node: src_node = write_edge.src src_node_connector = write_edge.src_conn src_subset = write_edge.data.src_subset diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index 3b5cff5ad4..e926cc25a5 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -9,6 +9,7 @@ from __future__ import annotations import abc +from collections import Counter from typing import TYPE_CHECKING, Iterable, Optional, Protocol import dace @@ -95,6 +96,7 @@ def _create_field_operator_impl( output_edge: gtir_dataflow.DataflowOutputEdge, output_type: ts.FieldType, map_exit: dace_nodes.MapExit, + output_consumer_count: dict[dace_nodes.AccessNode, int], ) -> gtir_to_sdfg_types.FieldopData: """ Helper method to allocate a temporary array that stores one field computed @@ -166,8 +168,11 @@ def _create_field_operator_impl( ) field_node = ctx.state.add_access(field_name) - # and here the edge writing the dataflow result data through the map exit node - output_edge.connect(map_exit, field_node, field_subset) + # and here the edge writing the dataflow result data through the map exit node. + # Note that we cannot remove the output data access node only if this is used + # for mutiple fields in a return tuple. + allow_removal_of_last_node = output_consumer_count[output_edge.result.dc_node] == 1 + output_edge.connect(map_exit, field_node, field_subset, allow_removal_of_last_node) return gtir_to_sdfg_types.FieldopData( field_node, ts.FieldType(field_dims, output_edge.result.gt_dtype), tuple(field_origin) @@ -217,17 +222,24 @@ def _create_field_operator( for edge in input_edges: edge.connect(map_entry) + # The same output node could be used for multiple fields in case of tuple return. + # In this case, the output access node cannot be removed. + consumer_count = Counter( + oedge.result.dc_node + for oedge in gtx_utils.flatten_nested_tuple((output_tree,)) + if oedge is not None + ) if isinstance(node_type, ts.FieldType): assert isinstance(output_tree, gtir_dataflow.DataflowOutputEdge) return _create_field_operator_impl( - ctx, sdfg_builder, domain, output_tree, node_type, map_exit + ctx, sdfg_builder, domain, output_tree, node_type, map_exit, consumer_count ) else: # handle tuples of fields output_symbol_tree = gtir_to_sdfg_utils.make_symbol_tree("x", node_type) return gtx_utils.tree_map( lambda output_edge, output_sym: _create_field_operator_impl( - ctx, sdfg_builder, domain, output_edge, output_sym.type, map_exit + ctx, sdfg_builder, domain, output_edge, output_sym.type, map_exit, consumer_count ) )(output_tree, output_symbol_tree) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index f315ce221f..5098b70ac9 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -22,6 +22,7 @@ from __future__ import annotations +from collections import Counter from typing import Iterable, Sequence import dace @@ -73,6 +74,7 @@ def _create_scan_field_operator_impl( output_domain: infer_domain.NonTupleDomainAccess, output_type: ts.FieldType, map_exit: dace_nodes.MapExit | None, + output_consumer_count: dict[dace_nodes.AccessNode, int], ) -> gtir_to_sdfg_types.FieldopData | None: """ Helper method to allocate a temporary array that stores one field computed @@ -146,7 +148,12 @@ def _create_scan_field_operator_impl( # Up to now the nested SDFG is writing into a transient data container that # has the size to hold one column. The function below, that does the connection, # will remove that transient and write directly to the result field. - inner_map_output_temporary_removed = output_edge.connect(map_exit, field_node, field_subset) + # Note that we cannot remove the output data access node only if this is used + # for mutiple fields in a return tuple. + allow_removal_of_last_node = output_consumer_count[output_edge.result.dc_node] == 1 + inner_map_output_temporary_removed = output_edge.connect( + map_exit, field_node, field_subset, allow_removal_of_last_node + ) if not inner_map_output_temporary_removed: raise ValueError("The scan nested SDFG is expected to write directly to the result field.") @@ -244,6 +251,14 @@ def _create_scan_field_operator( else im.sym("__gtir_unused_dummy_var", node_type) ) + # The same output node could be used for multiple fields in case of tuple return. + # In this case, the output access node cannot be removed. + consumer_count = Counter( + oedge.result.dc_node + for oedge in gtx_utils.flatten_nested_tuple((output,)) + if oedge is not None + ) + return gtx_utils.tree_map( lambda edge, domain, sym: _create_scan_field_operator_impl( ctx, @@ -252,6 +267,7 @@ def _create_scan_field_operator( domain, sym.type, map_exit, + consumer_count, ) )(output, output_domain, dummy_output_symbol) From 7a0a792e9fed88c5fb75cf8702b13cf22e32654d Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Fri, 15 May 2026 13:38:29 +0200 Subject: [PATCH 43/49] fix validation flag --- .../runners/dace/transformations/auto_optimize.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/auto_optimize.py b/src/gt4py/next/program_processors/runners/dace/transformations/auto_optimize.py index 799e8ad228..3f803ffcf7 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/auto_optimize.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/auto_optimize.py @@ -751,8 +751,8 @@ def _gt_auto_process_dataflow_inside_maps( # NestedSDFGs inside the ConditionalBlocks it fuses. sdfg.apply_transformations_repeated( gtx_transformations.FuseHorizontalConditionBlocks(), - validate=True, - validate_all=True, + validate=False, + validate_all=validate_all, ) # Move dataflow into the branches of the `if` such that they are only evaluated From 8e2079deb60c2c181ac71ef593f7dce9beda18f7 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Fri, 29 May 2026 14:05:02 +0200 Subject: [PATCH 44/49] fix[next-dace]: Reject MoveDataflowIntoIfBody when a tasklet free symbol clashes with an inner data descriptor A tasklet being relocated may reference a symbol from the outer SDFG as a free variable (no input connector). If the nested SDFG already contains a data descriptor with the same name, the relocated tasklet would silently read that inner descriptor instead of the outer symbol, producing wrong results. Fix _check_for_data_and_symbol_conflicts to return False whenever any symbol in required_symbols shares its name with a data descriptor in the nested SDFG. A regression test is included. Co-Authored-By: Claude Sonnet 4.6 --- .../move_dataflow_into_if_body.py | 7 ++ .../test_move_dataflow_into_if_body.py | 109 ++++++++++++++++++ 2 files changed, 116 insertions(+) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/move_dataflow_into_if_body.py b/src/gt4py/next/program_processors/runners/dace/transformations/move_dataflow_into_if_body.py index d7f7a36b96..528c95198e 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/move_dataflow_into_if_body.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/move_dataflow_into_if_body.py @@ -646,6 +646,13 @@ def _check_for_data_and_symbol_conflicts( if conflicting_symbol != str(symbol_mapping[conflicting_symbol]): return False + # A required symbol must not share its name with a data descriptor already present + # in the nested SDFG. If it did, after relocation the tasklet code would silently + # read the inner data descriptor instead of the outer symbol. + inner_arrays = if_block.sdfg.arrays + if any(sym in inner_arrays for sym in required_symbols): + return False + return True def _has_if_block_relocatable_dataflow( diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_move_dataflow_into_if_body.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_move_dataflow_into_if_body.py index f5b93da9ee..2240e45374 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_move_dataflow_into_if_body.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_move_dataflow_into_if_body.py @@ -2161,3 +2161,112 @@ def _make_outer_sdfg( # However, we would need to perform some more modifications. assert len([ie for ie in state.in_edges(nsdfg) if ie.data.data == "b"]) == 2 assert len([ac for ac in inner_ac if ac.data == "b"]) == 1 + + +def test_if_mover_symbol_clashes_with_inner_data(): + """Tests that relocation is rejected when a tasklet free symbol clashes + with a data descriptor inside the nested SDFG. + + The tasklet ``__out = my_var`` uses ``my_var`` as a free variable that refers + to the outer SDFG symbol. The nested SDFG also has ``my_var`` as an input + scalar data descriptor. If the transformation relocated the tasklet, the + free variable reference would silently resolve to the inner data descriptor + instead of the outer symbol, producing incorrect results. The transformation + must detect this conflict via ``_check_for_data_and_symbol_conflicts`` and + refuse to apply. + """ + sdfg = dace.SDFG(gtx_transformations.utils.unique_name("symbol_data_clash")) + state = sdfg.add_state(is_start_block=True) + + # `my_var` is a symbol in the outer SDFG, not a data container. + sdfg.add_symbol("my_var", dace.float64) + + for name in ["b", "out"]: + sdfg.add_array(name, shape=(10,), dtype=dace.float64, transient=False) + sdfg.add_scalar("val", dtype=dace.float64, transient=True) + sdfg.add_scalar("cond", dtype=dace.bool_, transient=True) + + me, mx = state.add_map("map", ndrange={"__i": "0:10"}) + + # Build the nested SDFG. It contains `my_var` as a non-transient input scalar + # data descriptor — the same name as the symbol in the outer SDFG. This + # creates the name conflict that the transformation must detect. + # The two input connectors are named `my_var` (true branch) and `other_var` + # (false branch). + inner_sdfg = dace.SDFG("if_body_with_my_var") + for name in ["my_var", "other_var", "__output"]: + inner_sdfg.add_scalar(name, dtype=dace.float64, transient=False) + inner_sdfg.add_scalar("__cond", dtype=dace.bool_, transient=False) + + if_region = dace.sdfg.state.ConditionalBlock("if_region") + inner_sdfg.add_node(if_region, is_start_block=True) + + then_body = dace.sdfg.state.ControlFlowRegion("then_body", sdfg=inner_sdfg) + tstate = then_body.add_state("true_branch", is_start_block=True) + tstate.add_nedge( + tstate.add_access("my_var"), + tstate.add_access("__output"), + dace.Memlet("my_var[0] -> [0]"), + ) + + else_body = dace.sdfg.state.ControlFlowRegion("else_body", sdfg=inner_sdfg) + fstate = else_body.add_state("false_branch", is_start_block=True) + fstate.add_nedge( + fstate.add_access("other_var"), + fstate.add_access("__output"), + dace.Memlet("other_var[0] -> [0]"), + ) + + if_region.add_branch(dace.sdfg.state.CodeBlock("__cond"), then_body) + if_region.add_branch(dace.sdfg.state.CodeBlock("not __cond"), else_body) + + if_block = state.add_nested_sdfg( + sdfg=inner_sdfg, + inputs={"my_var", "other_var", "__cond"}, + outputs={"__output"}, + ) + + # Tasklet with no input connectors that reads the outer symbol `my_var` as a + # free variable. This is the node the transformation would try to relocate. + tlet_sym = state.add_tasklet( # noqa: F841 [only needed to verify can_be_applied] + "symbol_to_scalar", + inputs={}, + outputs={"__out"}, + code="__out = my_var", + ) + val_ac = state.add_access("val") + cond_ac = state.add_access("cond") + + # Place `tlet_sym` inside the map via an empty memlet from `me`. + state.add_nedge(me, tlet_sym, dace.Memlet()) + state.add_edge(tlet_sym, "__out", val_ac, None, dace.Memlet("val[0]")) + state.add_edge(val_ac, None, if_block, "my_var", dace.Memlet("val[0]")) + + # `other_var` is a direct pass-through from `b` via the map. + state.add_edge(state.add_access("b"), None, me, "IN_b", dace.Memlet("b[0:10]")) + state.add_edge(me, "OUT_b", if_block, "other_var", dace.Memlet("b[__i]")) + me.add_scope_connectors("b") + + # Condition: a tasklet that writes True to `cond` when `__i` is even. + cond_tlet = state.add_tasklet( + "cond_tasklet", + inputs={}, + outputs={"__out"}, + code="__out = (__i % 2 == 0)", + ) + state.add_nedge(me, cond_tlet, dace.Memlet()) + state.add_edge(cond_tlet, "__out", cond_ac, None, dace.Memlet("cond[0]")) + state.add_edge(cond_ac, None, if_block, "__cond", dace.Memlet("cond[0]")) + + # Output path. + state.add_edge(if_block, "__output", mx, "IN_out", dace.Memlet("out[__i]")) + state.add_edge(mx, "OUT_out", state.add_access("out"), None, dace.Memlet("out[0:10]")) + mx.add_scope_connectors("out") + + sdfg.validate() + + # The transformation must NOT apply. The free symbol `my_var` in `tlet_sym` + # has the same name as the input data descriptor `my_var` inside the nested + # SDFG. Relocating `tlet_sym` would silently replace the outer-symbol + # reference with a read of the inner scalar, producing wrong results. + _perform_test(sdfg, expected_applies=0, if_block=if_block) From d36994be4261d5903e7eda29d6a3118c05bdd54a Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Thu, 28 May 2026 11:42:33 +0200 Subject: [PATCH 45/49] CI: Use production credentials for beverin (#2611) The CI on Beverin vCluster now requires production credentials. It was using TDS credentials before. --- ci/cscs-ci-ext-config.yml | 28 ---------------------------- ci/cscs-ci.yml | 5 ++--- 2 files changed, 2 insertions(+), 31 deletions(-) delete mode 100644 ci/cscs-ci-ext-config.yml diff --git a/ci/cscs-ci-ext-config.yml b/ci/cscs-ci-ext-config.yml deleted file mode 100644 index 6a35d26141..0000000000 --- a/ci/cscs-ci-ext-config.yml +++ /dev/null @@ -1,28 +0,0 @@ -# -# GT4Py - GridTools Framework -# -# Copyright (c) 2014-2024, ETH Zurich -# All rights reserved. -# -# Please, refer to the LICENSE file in the root directory. -# SPDX-License-Identifier: BSD-3-Clause -# - -include: - - remote: 'https://gitlab.com/cscs-ci/recipes/-/raw/master/templates/v2/.ci-ext.yml' - -# TDS needs different credentials -.tds-container-runner-beverin: - variables: - F7T_CLIENT_ID: $F7T_TDS_CLIENT_ID - F7T_CLIENT_SECRET: $F7T_TDS_CLIENT_SECRET - -.tds-container-runner-beverin-mi200: - extends: - - .container-runner-beverin-mi200 - - .tds-container-runner-beverin - -.tds-container-runner-beverin-mi300: - extends: - - .container-runner-beverin-mi300 - - .tds-container-runner-beverin diff --git a/ci/cscs-ci.yml b/ci/cscs-ci.yml index 0f4475ad61..955257e8ac 100644 --- a/ci/cscs-ci.yml +++ b/ci/cscs-ci.yml @@ -10,7 +10,6 @@ include: - remote: 'https://gitlab.com/cscs-ci/recipes/-/raw/master/templates/v2/.ci-ext.yml' - - local: 'ci/cscs-ci-ext-config.yml' # Note: # block-name-with-dashes -> defined in remote cscs-ci ext include @@ -48,7 +47,7 @@ stages: # To override the defaults, just define these variables in the actual job. DOCKER_BUILD_ARGS: '["BASE_IMAGE", "CACHE_DIR", "EXTRA_APTGET", "EXTRA_UV_ENV_VARS", "EXTRA_UV_PIP_ARGS", "EXTRA_UV_SYNC_ARGS", "PY_VERSION", "UV_VERSION", "WORKDIR_PATH" ]' PERSIST_IMAGE_NAME: ${CSCS_REGISTRY_PATH}/public/${ARCH}/base/gt4py-ci-${PY_VERSION} # The $DOCKER_TAG tag is added in the before_script of .dynamic-image-name - WATCH_FILECHANGES: 'ci/Dockerfile ci/cscs-ci.yml ci/cscs-ci-ext-config.yml uv.lock' + WATCH_FILECHANGES: 'ci/Dockerfile ci/cscs-ci.yml uv.lock' parallel: matrix: - PY_VERSION: *test_python_versions @@ -154,7 +153,7 @@ test_cscs_gh200: test_cscs_amd_rocm: extends: - - .tds-container-runner-beverin-mi200 + - .container-runner-beverin-mi200 - .test_common needs: - build_cscs_amd_rocm From 326f78c1da4e7201be4eb587e7640e64c5ce6eb2 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Mon, 1 Jun 2026 13:34:19 +0200 Subject: [PATCH 46/49] update ci config --- ci/cscs-ci.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ci/cscs-ci.yml b/ci/cscs-ci.yml index 955257e8ac..7e613a55a5 100644 --- a/ci/cscs-ci.yml +++ b/ci/cscs-ci.yml @@ -162,11 +162,11 @@ test_cscs_amd_rocm: SLURM_GPUS_PER_NODE: 4 GT4PY_BUILD_JOBS: 8 # Limit test parallelism to avoid "OSError: too many open files" in the gt4py build stage. - PYTEST_XDIST_AUTO_NUM_WORKERS: 32 + PYTEST_XDIST_AUTO_NUM_WORKERS: 16 SLURM_PARTITION: mi300 CMAKE_PREFIX_PATH: /opt/rocm # for next CUDA_HOME: /opt/rocm # for cartesian - SLURM_TIMELIMIT: 20 # relaxed relative to gh200 as there is no pressure on the queue + SLURM_TIMELIMIT: 30 # relaxed relative to gh200 as there is no pressure on the queue rules: - *exclude_variants_rules - if: $SUBPACKAGE == 'cartesian' && $VARIANT == 'internal' && $SUBVARIANT == 'cpu' From caac94839bea05d853eea1d5f3975186b90d2e4b Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Sat, 19 Sep 2026 08:21:46 +0200 Subject: [PATCH 47/49] feat[next-dace]: support constant lists from make_const_list --- .../runners/dace/lowering/gtir_to_sdfg.py | 26 ++++++++++- .../dace/lowering/gtir_to_sdfg_lambda.py | 2 +- .../dace/lowering/gtir_to_sdfg_utils.py | 4 ++ .../runners/dace/workflow/translation.py | 43 +++++++++++++++++-- 4 files changed, 68 insertions(+), 7 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index e645068f64..9b52e7f21e 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -571,8 +571,20 @@ def make_field( ) -> gtir_to_sdfg_types.FieldopData: local_dims = [dim for dim in data_type.dims if dim.kind == gtx_common.DimensionKind.LOCAL] if len(local_dims) == 0: - # do nothing: the field domain consists of all global dimensions - field_type = data_type + if isinstance(data_type.dtype, ts.ListType) and data_type.dtype.offset_type is None: + # A field of constant lists (as produced by 'make_const_list') carries a + # list dtype without offset provider. We tag it with the magic local + # dimension so that it is handled like any other neighbor-list field. + field_type = ts.FieldType( + dims=data_type.dims, + dtype=ts.ListType( + element_type=data_type.dtype.element_type, + offset_type=gtir_to_sdfg_utils.CONST_DIM, + ), + ) + else: + # do nothing: the field domain consists of all global dimensions + field_type = data_type elif len(local_dims) == 1: local_dim = local_dims[0] # the local dimension is converted into `ListType` data element @@ -912,6 +924,16 @@ def _add_storage( elif isinstance(gt_type, ts.FieldType): if len(gt_type.dims) == 0: + if isinstance(gt_type.dtype, ts.ListType): + # A zero-dimensional field with list dtype represents a field of + # constant lists (as produced by 'make_const_list'), which holds a + # single value broadcast over the local dimension. We store it as a + # single-element 1D array, consistent with the representation of + # constant-list value expressions (see '_make_value' with 'use_array'). + assert isinstance(gt_type.dtype.element_type, ts.ScalarType) + dc_dtype = gtx_dace_args.as_dace_type(gt_type.dtype.element_type) + sdfg.add_array(name, (1,), dc_dtype, transient=transient) + return [(name, gt_type)] # represent zero-dimensional fields as scalar arguments return self._add_storage( sdfg=sdfg, diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 654e1fc9ff..3ff954c219 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -59,7 +59,7 @@ # Magic local dimension used for list of values with length known at compile-time. -_CONST_DIM: Final = gtx_common.Dimension(value="_CONST_DIM", kind=gtx_common.DimensionKind.LOCAL) +_CONST_DIM: Final = gtir_to_sdfg_utils.CONST_DIM @dataclasses.dataclass(frozen=True) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py index a6f591c04c..f841036e9e 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py @@ -24,6 +24,10 @@ _TASKLET_CONNECTOR_PREFIX: Final[str] = "__tlet_" """Prefix string to be used for tasklet connectors.""" +CONST_DIM: Final = gtx_common.Dimension(value="_CONST_DIM", kind=gtx_common.DimensionKind.LOCAL) +"""Magic local dimension used for a list of values with length known at compile-time, +as produced by 'make_const_list'.""" + def debug_info( node: gtir.Node, *, default: Optional[dace.dtypes.DebugInfo] = None diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index d4a78added..36b11e4ee3 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -10,6 +10,7 @@ import dataclasses import functools +from collections.abc import Iterable, Iterator from typing import Any, Optional import dace @@ -19,6 +20,7 @@ from gt4py.next import common from gt4py.next.instrumentation import metrics from gt4py.next.iterator import ir as itir +from gt4py.next.iterator.ir_utils import common_pattern_matcher as cpm from gt4py.next.iterator.transforms import pass_manager from gt4py.next.otf import artifacts, stages, workflow from gt4py.next.otf.binding import interface @@ -31,6 +33,33 @@ from gt4py.next.type_system import type_specifications as ts +def _is_neighbors_or_lifted_neighbors(arg: itir.Expr) -> bool: + """Whether `arg` is a `neighbors` call or a lift wrapping one (transitively).""" + if cpm.is_call_to(arg, "neighbors"): + return True + return cpm.is_applied_lift(arg) and any( + _is_neighbors_or_lifted_neighbors(nested_arg) for nested_arg in arg.args + ) + + +def _has_neighbors_argument(reduce_node: itir.FunCall) -> bool: + """Whether an applied `reduce` operates directly on a (lifted) `neighbors` argument. + + Such a reduction is not lowered to SDFG natively and must be unrolled. A reduction + over a materialized neighbor-list field (e.g. the result of a `concat_where`) has a + plain iterator argument instead and is excluded here, since it is lowered natively. + """ + + def flatten(args: Iterable[itir.Expr]) -> Iterator[itir.Expr]: + for arg in args: + if cpm.is_call_to(arg, "if_"): + yield from flatten(arg.args[1:3]) + else: + yield arg + + return any(_is_neighbors_or_lifted_neighbors(arg) for arg in flatten(reduce_node.args)) + + def find_constant_symbols( ir: itir.Program, sdfg: dace.SDFG, @@ -389,11 +418,17 @@ def _preprocess_program( new_program = apply_common_transforms(program, unroll_reduce=False) if any( - node.id == "neighbors" - for node in new_program.pre_walk_values().if_isinstance(itir.SymRef) + cpm.is_applied_reduce(node) and _has_neighbors_argument(node) + for node in new_program.pre_walk_values().if_isinstance(itir.FunCall) ): - # if we don't unroll, there may be lifts left in the itir which can't - # be lowered to SDFG. In this case, just retry with unrolled reductions. + # A `reduce` applied directly to a `neighbors` (or lifted `neighbors`) + # argument cannot be lowered to SDFG as-is, so we retry with unrolled + # reductions (which also inlines any lifts left in the itir). A `reduce` + # over a materialized neighbor-list field, e.g. the result of a + # `concat_where`, is lowered natively and must not trigger this path: + # unrolling such a reduction is either unnecessary or, on meshes with + # skip values, outright fails since there is no `neighbors` iterator from + # which to build the skip-value check. new_program = apply_common_transforms(program, unroll_reduce=True) return new_program From da80b9c9cd8bdcc8039075dc995c707ef79d56b2 Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Sat, 19 Sep 2026 22:50:11 +0200 Subject: [PATCH 48/49] fix[next-dace]: unroll reductions when lifts remain The mixed-IR DaCe path only inlines remaining lifts via the unroll-reduce retry loop. The previous retry condition fired only for a reduce applied directly to a (lifted) neighbors argument, so a lift-only program (e.g. a sparse field access lowered to neighbors over an identity lift, without any reduction) kept the lift and failed lowering with NotImplementedError. Retry with unrolled reductions whenever an applied lift is left in the itir, in addition to the reduce-over-neighbors case. A reduce over a materialized neighbor-list field (e.g. a concat_where result) still matches neither and is lowered natively. --- .../runners/dace/workflow/translation.py | 21 +++++++++++-------- 1 file changed, 12 insertions(+), 9 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index 36b11e4ee3..46b74ce20e 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -418,17 +418,20 @@ def _preprocess_program( new_program = apply_common_transforms(program, unroll_reduce=False) if any( - cpm.is_applied_reduce(node) and _has_neighbors_argument(node) + cpm.is_applied_lift(node) + or (cpm.is_applied_reduce(node) and _has_neighbors_argument(node)) for node in new_program.pre_walk_values().if_isinstance(itir.FunCall) ): - # A `reduce` applied directly to a `neighbors` (or lifted `neighbors`) - # argument cannot be lowered to SDFG as-is, so we retry with unrolled - # reductions (which also inlines any lifts left in the itir). A `reduce` - # over a materialized neighbor-list field, e.g. the result of a - # `concat_where`, is lowered natively and must not trigger this path: - # unrolling such a reduction is either unnecessary or, on meshes with - # skip values, outright fails since there is no `neighbors` iterator from - # which to build the skip-value check. + # We retry with unrolled reductions (whose fixed-point loop also inlines + # the remaining lifts) in two cases that the SDFG lowering cannot handle + # as-is: + # - an applied `lift` is left in the itir, or + # - a `reduce` is applied directly to a (lifted) `neighbors` argument. + # A `reduce` over a materialized neighbor-list field, e.g. the result of a + # `concat_where`, does not match either condition: it is lowered natively + # and must not trigger this path, since unrolling such a reduction is + # unnecessary and, on meshes with skip values, outright fails (there is no + # `neighbors` iterator from which to build the skip-value check). new_program = apply_common_transforms(program, unroll_reduce=True) return new_program From 73973f6108749baa6c6da7763f035ff4c78d30ae Mon Sep 17 00:00:00 2001 From: Edoardo Paone Date: Sat, 19 Sep 2026 22:52:28 +0200 Subject: [PATCH 49/49] fix test --- .../runners_tests/dace_tests/test_dace_translation.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index e3cf79727a..52ac5b46d9 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -484,6 +484,7 @@ def test_translation_source_code_invariant_under_guid_change(): translator = dace_wf_translation.DaCeTranslator( device_type=core_defs.DeviceType.CPU, + apply_common_transform=False, auto_optimize=False, auto_optimize_args=None, async_sdfg_call=False,