diff --git a/.cscs-ci/default.yml b/.cscs-ci/default.yml index d19c6fb6f7..c5494e4f91 100644 --- a/.cscs-ci/default.yml +++ b/.cscs-ci/default.yml @@ -120,11 +120,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' diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index 067dde468f..ff8d110bfa 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -139,6 +139,7 @@ def apply_common_transforms( unroll_reduce=False, common_subexpression_elimination=True, force_inline_lambda_args=False, + 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, itir.Expr]] = None, @@ -192,7 +193,8 @@ def apply_common_transforms( ) 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 @@ -264,6 +266,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/__init__.py b/src/gt4py/next/program_processors/runners/dace/__init__.py index 0e560fa761..4fc393f663 100644 --- a/src/gt4py/next/program_processors/runners/dace/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/__init__.py @@ -11,6 +11,8 @@ from gt4py.next.program_processors.runners.dace.workflow.backend import ( make_dace_backend, run_dace_cpu, + run_dace_cpu_gt, + run_dace_cpu_gt_noopt, run_dace_cpu_noopt, run_dace_gpu, run_dace_gpu_noopt, @@ -21,6 +23,8 @@ "get_sdfg_args", "make_dace_backend", "run_dace_cpu", + "run_dace_cpu_gt", + "run_dace_cpu_gt_noopt", "run_dace_cpu_noopt", "run_dace_gpu", "run_dace_gpu_noopt", 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 f655de3d6a..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) @@ -283,6 +283,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`. @@ -294,24 +295,27 @@ 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): - # 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): - # 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 @@ -435,6 +439,7 @@ class LambdaToDataflow(eve.NodeVisitor): def __post_init__(self) -> None: builtin_dispatch: dict[str, Callable[[gtir.FunCall], VisitResult]] = { + "can_deref": self._visit_can_deref, "deref": self._visit_deref, "if_": self._visit_if, "neighbors": self._visit_neighbors, @@ -598,6 +603,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. @@ -1227,7 +1277,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): @@ -1573,7 +1623,7 @@ def _make_cartesian_shift( if isinstance(index_expr, SymbolExpr) and isinstance(offset_expr, SymbolExpr): # purely symbolic expression which can be interpreted at compile time new_index = SymbolExpr( - index_expr.value + offset_expr.value, + dace.symbolic.pystr_to_symbolic(f"{index_expr.value} + {offset_expr.value}"), index_expr.dc_dtype, ) else: @@ -1640,9 +1690,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. @@ -1651,33 +1701,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( @@ -1692,17 +1765,14 @@ 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] - assert isinstance(conn_source_index, SymbolExpr) shifted_indices = { dim: idx for dim, idx in it.indices.items() if dim != conn_type.source_dim } - if isinstance(offset_expr, SymbolExpr): + if isinstance(offset_expr, SymbolExpr) and isinstance(conn_source_index, SymbolExpr): # use memlet to retrieve the neighbor index shifted_indices[conn_type.codomain] = MemletExpr( dc_node=conn_node, @@ -1977,6 +2047,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 8e14ae41bd..4eb0de88d5 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,11 +9,13 @@ from __future__ import annotations import abc +from collections import Counter from typing import TYPE_CHECKING, Iterable, Optional, Protocol import dace from dace import nodes as dace_nodes, 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 ( @@ -69,16 +71,18 @@ def _parse_fieldop_arg( ctx: gtir_to_sdfg.SubgraphContext, sdfg_builder: gtir_to_sdfg.SDFGBuilder, domain: gtir_domain.FieldopDomain, -) -> gtir_to_sdfg_lambda.IteratorExpr | gtir_to_sdfg_lambda.MemletExpr: +) -> MaybeNestedInTuple[gtir_to_sdfg_lambda.IteratorExpr | gtir_to_sdfg_lambda.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, ctx.sdfg))(arg) def _create_field_operator_impl( @@ -88,6 +92,7 @@ def _create_field_operator_impl( output_edge: gtir_to_sdfg_lambda.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 @@ -158,8 +163,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) @@ -169,10 +177,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_to_sdfg_lambda.DataflowInputEdge], - output_edge: gtir_to_sdfg_lambda.DataflowOutputEdge, + output_tree: MaybeNestedInTuple[gtir_to_sdfg_lambda.DataflowOutputEdge], ) -> gtir_to_sdfg_types.FieldopResult: """ Helper method to build the output of a field operator. @@ -183,11 +191,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: @@ -209,7 +217,26 @@ 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) + # 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_to_sdfg_lambda.DataflowOutputEdge) + return _create_field_operator_impl( + 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, consumer_count + ) + )(output_tree, output_symbol_tree) def translate_as_fieldop( @@ -243,9 +270,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( @@ -260,7 +284,7 @@ def translate_as_fieldop( # the input value (a scalar, a field slice or a field of constant lists) # on the output field type. 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. @@ -280,13 +304,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_to_sdfg_lambda.translate_lambda_to_dataflow( + input_edges, output_edges = gtir_to_sdfg_lambda.translate_lambda_to_dataflow( ctx.sdfg, ctx.state, sdfg_builder, stencil_expr, fieldop_args ) - assert isinstance(output_edge, gtir_to_sdfg_lambda.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 ) 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 bbd53c2b9d..9ebc7c232e 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 @@ -23,6 +23,7 @@ from __future__ import annotations import copy +from collections import Counter from typing import Iterable, Sequence import dace @@ -49,43 +50,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_to_sdfg_lambda.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_to_sdfg_lambda.IteratorExpr | gtir_to_sdfg_lambda.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_to_sdfg_lambda.MemletExpr: - arg_expr = arg.get_local_view(field_domain, ctx.sdfg) - if isinstance(arg_expr, gtir_to_sdfg_lambda.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_to_sdfg_lambda.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( @@ -95,6 +76,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 @@ -168,7 +150,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.") @@ -266,6 +253,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, @@ -274,6 +269,7 @@ def _create_scan_field_operator( domain, sym.type, map_exit, + consumer_count, ) )(output, output_domain, dummy_output_symbol) @@ -429,7 +425,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 ] # still inside the 'compute' state, generate the dataflow representing the stencil 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/transformations/strides.py b/src/gt4py/next/program_processors/runners/dace/transformations/strides.py index ca5be147c7..385980b109 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,7 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import warnings from typing import Optional, TypeAlias import dace @@ -528,7 +529,11 @@ 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 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 # 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 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 c0ea33daf0..17377c5f94 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py @@ -79,6 +79,7 @@ class Params: def make_dace_backend( gpu: bool, + apply_common_transform: bool = False, auto_optimize: bool = True, async_sdfg_call: bool = True, optimization_args: dict[str, Any] | None = None, @@ -92,6 +93,8 @@ def make_dace_backend( Args: gpu: Enable GPU transformations and code generation. + 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. @@ -158,6 +161,7 @@ def make_dace_backend( gpu=gpu, auto_optimize=auto_optimize, external_workspace=external_workspace, + 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, @@ -188,3 +192,15 @@ def make_dace_backend( auto_optimize=False, async_sdfg_call=True, ) +run_dace_cpu_gt = make_dace_backend( + gpu=False, + apply_common_transform=True, + auto_optimize=True, + async_sdfg_call=False, +) +run_dace_cpu_gt_noopt = make_dace_backend( + gpu=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 f26e982eff..46b74ce20e 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,8 @@ from __future__ import annotations import dataclasses +import functools +from collections.abc import Iterable, Iterator from typing import Any, Optional import dace @@ -17,7 +19,9 @@ from gt4py._core import definitions as core_defs from gt4py.next import common 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.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 from gt4py.next.program_processors.runners.dace import ( @@ -29,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, @@ -352,6 +383,7 @@ class DaCeTranslator( ], ): device_type: core_defs.DeviceType + apply_common_transform: bool auto_optimize: bool auto_optimize_args: dict[str, Any] | None async_sdfg_call: bool @@ -370,6 +402,40 @@ 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, + 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, + ) + + new_program = apply_common_transforms(program, unroll_reduce=False) + + if any( + 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) + ): + # 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 + def _generate_sdfg_without_configuring_dace( self, ir: itir.Program, @@ -377,11 +443,14 @@ 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, - use_max_domain_range_on_unstructured_shift=self.use_max_domain_range_on_unstructured_shift, - offset_provider=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 a75f3e08ff..8cbe56d713 100644 --- a/tests/next_tests/definitions.py +++ b/tests/next_tests/definitions.py @@ -72,6 +72,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): @@ -213,6 +215,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_GPU: GTFN_SKIP_TEST_LIST diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 18617a6f91..e793a4cd1a 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -114,6 +114,14 @@ def __gt_allocator__( next_tests.definitions.OptionalProgramBackendId.DACE_CPU_NO_OPT, marks=pytest.mark.uses_dace, ), + pytest.param( + next_tests.definitions.OptionalProgramBackendId.DACE_CPU_GT, + marks=pytest.mark.uses_dace, + ), + pytest.param( + next_tests.definitions.OptionalProgramBackendId.DACE_CPU_GT_NO_OPT, + marks=pytest.mark.uses_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 879094f123..d274b56e99 100644 --- a/tests/next_tests/unit_tests/conftest.py +++ b/tests/next_tests/unit_tests/conftest.py @@ -64,6 +64,10 @@ def _program_processor(request) -> tuple[ProgramProcessor, bool]: (next_tests.definitions.OptionalProgramBackendId.DACE_CPU_NO_OPT, True), marks=pytest.mark.uses_dace, ), + pytest.param( + (next_tests.definitions.OptionalProgramBackendId.DACE_CPU_GT_NO_OPT, True), + marks=pytest.mark.uses_dace, + ), ], ids=lambda p: p[0].short_id() if p[0] is not None else "None", ) 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 87e212afc9..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 @@ -69,6 +69,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, @@ -483,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,