diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/__init__.py b/src/gt4py/next/program_processors/runners/dace/transformations/__init__.py index a5289922b6..b41ba3269c 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/__init__.py @@ -85,6 +85,7 @@ gt_propagate_strides_from_access_node, gt_propagate_strides_of, ) +from .trivial_map_dimension_folding import TrivialMapDimensionFolding from .utils import gt_configure_transient_lifetime from .write_back_buffer_elimination import GT4PyWriteBackBufferElimination @@ -123,6 +124,7 @@ "SplitAccessNode", "SplitConsumerMemlet", "TransientMemoryMode", + "TrivialMapDimensionFolding", "VerticalMapFusionCallback", "VerticalMapSplitCallback", "constants", 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 2766e58683..5091dd4283 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 @@ -523,6 +523,15 @@ def _gt_auto_process_top_level_maps( # has been solved. vertical_map_fusion._single_use_data = single_use_data + # Fold single iteration Map dimensions first, otherwise their Memlets still + # refer to the parameter symbolically and `Range.covers()` fails to see that + # a producer covers a consumer, rejecting legal fusions. + sdfg.apply_transformations_repeated( + gtx_transformations.TrivialMapDimensionFolding(only_toplevel_maps=True), + validate=False, + validate_all=validate_all, + ) + sdfg.apply_transformations_repeated( vertical_map_fusion, validate=False, diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/trivial_map_dimension_folding.py b/src/gt4py/next/program_processors/runners/dace/transformations/trivial_map_dimension_folding.py new file mode 100644 index 0000000000..9c3542e6d9 --- /dev/null +++ b/src/gt4py/next/program_processors/runners/dace/transformations/trivial_map_dimension_folding.py @@ -0,0 +1,103 @@ +# 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 + +from __future__ import annotations + +import copy +from typing import Any, Optional, Union + +import dace +from dace import properties as dace_properties, transformation as dace_transformation +from dace.sdfg import nodes as dace_nodes + + +@dace_properties.make_properties +class TrivialMapDimensionFolding(dace_transformation.SingleStateTransformation): + """Replaces the parameter of a single iteration Map dimension by its value. + + A Map dimension such as `i = c:c+1:1` can only ever take the value `c`, but the + Memlets inside its scope still refer to `i` symbolically. `MapFusionVertical` + compares producer and consumer subsets with `Range.covers()`, which does not + know the Map range, so a producer writing `a[c]` is not recognized as covering a + consumer reading `a[i - o]` even though both denote the same element, and a legal + fusion is rejected. Folding the value into the Memlets removes that blind spot. + + Note: + Unlike DaCe's native `TrivialMapElimination` the dimension is not removed, + only the single value is folded into the body of the Map. + + Args: + only_toplevel_maps: Only process Maps that are on the top level. + """ + + map_entry = dace_transformation.PatternNode(dace_nodes.MapEntry) + + only_toplevel_maps = dace_properties.Property( + dtype=bool, + default=False, + desc="See docs.", + ) + + def __init__(self, only_toplevel_maps: Optional[bool] = None, **kwargs: Any) -> None: + super().__init__(**kwargs) + if only_toplevel_maps is not None: + self.only_toplevel_maps = only_toplevel_maps + + @classmethod + def expressions(cls) -> Any: + return [dace.sdfg.utils.node_path_graph(cls.map_entry)] + + @staticmethod + def _single_iteration_parameters(map_: dace_nodes.Map) -> dict[str, Any]: + """Returns the parameters of `map_` that can only take a single value.""" + return { + param: rng[0] + for param, rng in zip(map_.params, map_.range.ranges) + if (rng[0] == rng[1]) == True # noqa: E712 [true-false-comparison] # SymPy comparison + } + + def can_be_applied( + self, + graph: Union[dace.SDFGState, dace.SDFG], + expr_index: int, + sdfg: dace.SDFG, + permissive: bool = False, + ) -> bool: + map_entry: dace_nodes.MapEntry = self.map_entry + if self.only_toplevel_maps and graph.entry_node(map_entry) is not None: + return False + + replacements = self._single_iteration_parameters(map_entry.map) + if not replacements: + return False + + # Only apply if a parameter is still referenced, otherwise the transformation + # would apply again and again on the same Map. + scope = graph.scope_subgraph(map_entry, include_entry=True, include_exit=True) + for edge in scope.edges(): + if edge.data is not None and any( + str(sym) in replacements for sym in edge.data.free_symbols + ): + return True + for node in scope.nodes(): + if node is not map_entry and any(str(sym) in replacements for sym in node.free_symbols): + return True + return False + + def apply(self, graph: Union[dace.SDFGState, dace.SDFG], sdfg: dace.SDFG) -> None: + map_entry: dace_nodes.MapEntry = self.map_entry + replacements = self._single_iteration_parameters(map_entry.map) + scope = graph.scope_subgraph(map_entry, include_entry=True, include_exit=True) + + # `replace_dict()` would also rewrite the Map's own parameters and range, which + # would drop the dimension, so they are restored afterwards. + saved_params = copy.deepcopy(map_entry.map.params) + saved_range = copy.deepcopy(map_entry.map.range) + scope.replace_dict({param: str(value) for param, value in replacements.items()}) + map_entry.map.params = saved_params + map_entry.map.range = saved_range diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_trivial_map_dimension_folding.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_trivial_map_dimension_folding.py new file mode 100644 index 0000000000..d184c27aa6 --- /dev/null +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_trivial_map_dimension_folding.py @@ -0,0 +1,106 @@ +# 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 + +from __future__ import annotations + +import copy + +import numpy as np +import pytest + +dace = pytest.importorskip("dace") + +from gt4py.next.program_processors.runners.dace import ( + transformations as gtx_transformations, +) + +from . import util + + +def _mk_single_iteration_sdfg() -> dace.SDFG: + """A Map whose second dimension has exactly one iteration.""" + sdfg = dace.SDFG(util.unique_name("single_iteration_map")) + for name in ["a", "b"]: + sdfg.add_array(name, shape=(20, 10), dtype=dace.float64, transient=False) + + state = sdfg.add_state(is_start_block=True) + state.add_mapped_tasklet( + "comp", + map_ranges={"__i": "0:20", "__j": "7:8"}, + inputs={"__in": dace.Memlet("a[__i, __j]")}, + code="__out = __in + 1.0", + outputs={"__out": dace.Memlet("b[__i, __j]")}, + external_edges=True, + ) + sdfg.validate() + return sdfg + + +def test_trivial_map_dimension_folding(): + sdfg = _mk_single_iteration_sdfg() + + ref = { + name: np.array(np.random.rand(*desc.shape), copy=True, dtype=desc.dtype.as_numpy_dtype()) + for name, desc in sdfg.arrays.items() + if not desc.transient + } + res = copy.deepcopy(ref) + util.compile_and_run_sdfg(sdfg, **ref) + + # a single application although applied repeatedly, i.e. the rewrite is a fixpoint + nb_apply = sdfg.apply_transformations_repeated( + gtx_transformations.TrivialMapDimensionFolding(), + validate=True, + validate_all=True, + ) + assert nb_apply == 1 + + map_entry = next( + node + for state in sdfg.states() + for node in state.nodes() + if isinstance(node, dace.sdfg.nodes.MapEntry) + ) + # The dimension must survive, only its uses are rewritten. + assert map_entry.map.params == ["__i", "__j"] + assert str(map_entry.map.range[1][0]) == "7" + assert all( + "__j" not in str(subset) + for state in sdfg.states() + for edge in state.edges() + if edge.data is not None + for subset in (edge.data.subset, edge.data.other_subset) + if subset is not None + ) + + util.compile_and_run_sdfg(sdfg, **res) + assert all(np.allclose(ref[name], res[name]) for name in ref.keys()) + + +def test_trivial_map_dimension_folding_ignores_multi_iteration(): + """A Map without a single iteration dimension must not be touched.""" + sdfg = dace.SDFG(util.unique_name("multi_iteration_map")) + for name in ["a", "b"]: + sdfg.add_array(name, shape=(20, 10), dtype=dace.float64, transient=False) + state = sdfg.add_state(is_start_block=True) + state.add_mapped_tasklet( + "comp", + map_ranges={"__i": "0:20", "__j": "0:10"}, + inputs={"__in": dace.Memlet("a[__i, __j]")}, + code="__out = __in + 1.0", + outputs={"__out": dace.Memlet("b[__i, __j]")}, + external_edges=True, + ) + sdfg.validate() + + assert ( + sdfg.apply_transformations_repeated( + gtx_transformations.TrivialMapDimensionFolding(), validate=True + ) + == 0 + )