From 84c59784daeb9abe666e9e12bc3ef1a950fe1ae7 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Sun, 9 Aug 2026 22:02:50 +0200 Subject: [PATCH 1/3] feat[next-dace]: fold single iteration Map dimensions A Map dimension `i = c:c+1:1` can only ever take the value `c`, but the Memlets inside its scope keep referring 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]` and a legal fusion is rejected. The new transformation folds the value into the scope and runs before the top level fusion rounds. The dimension itself is kept, unlike DaCe's `TrivialMapElimination` which removes it and thereby prevents the Map from being scheduled on the GPU. On `compute_perturbed_quantities_and_interpolation` this collapses the three surface level kernels into one, 10 -> 9 kernels. It also removes a non-determinism: the same program compiled to 10 or 11 kernels depending on which other variants were compiled in the same session. Co-Authored-By: Claude Opus 5 --- .../runners/dace/transformations/__init__.py | 2 + .../dace/transformations/auto_optimize.py | 9 ++ .../trivial_map_dimension_folding.py | 109 +++++++++++++++++ .../test_trivial_map_dimension_folding.py | 114 ++++++++++++++++++ 4 files changed, 234 insertions(+) create mode 100644 src/gt4py/next/program_processors/runners/dace/transformations/trivial_map_dimension_folding.py create mode 100644 tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_trivial_map_dimension_folding.py 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 f8cb38334f..4405d1ef01 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 @@ -121,6 +122,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 2a6997406a..a999e0597c 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 @@ -519,6 +519,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..08473a6be3 --- /dev/null +++ b/src/gt4py/next/program_processors/runners/dace/transformations/trivial_map_dimension_folding.py @@ -0,0 +1,109 @@ +# 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: + The dimension itself is kept, in contrast to DaCe's `TrivialMapElimination` + which removes it. Removing it changes the shape of the Map and prevents it + from being scheduled on the GPU, whereas folding only rewrites the uses of + the parameter and leaves the iteration space, the schedule and any blocking + applied later untouched. + + 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="Only process Maps that are on the top level.", + ) + + 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 and (rng[2] == 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. + # NOTE: `include_entry` is needed because the uses are on the out edges of + # the MapEntry, which are not part of the scope subgraph otherwise. + 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()` 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) + for param, value in replacements.items(): + scope.replace(param, value) + 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..096ebfbf5e --- /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,114 @@ +# 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) + + 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(edge.data.subset) + for state in sdfg.states() + for edge in state.edges() + if edge.data is not None and edge.data.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_is_a_fixpoint(): + """A second application must not be possible, else the pass would not terminate.""" + sdfg = _mk_single_iteration_sdfg() + assert ( + sdfg.apply_transformations_repeated( + gtx_transformations.TrivialMapDimensionFolding(), validate=True + ) + == 1 + ) + + +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 + ) From bda617549828eb8fb97046fb2f752d949836d1e5 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Mon, 17 Aug 2026 09:32:42 +0200 Subject: [PATCH 2/3] refactor[next-dace]: address review comments on trivial Map dimension folding - drop the `rng[2] == 1` step check: with `start == end` the dimension has a single iteration for any step, so the step carries no information here. - fold with a single `replace_dict()` instead of one `replace()` per parameter. - shorten the note contrasting the transformation with `TrivialMapElimination` to what distinguishes them, the dimension being kept rather than removed. - test the `other_subset` of every Memlet as well, not only the `subset`. - drop the separate fixpoint test: the main test already applies the transformation repeatedly and asserts a single application. Co-Authored-By: Claude Opus 5 --- .../trivial_map_dimension_folding.py | 14 +++++--------- .../test_trivial_map_dimension_folding.py | 18 +++++------------- 2 files changed, 10 insertions(+), 22 deletions(-) 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 index 08473a6be3..111d243aec 100644 --- 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 @@ -28,11 +28,8 @@ class TrivialMapDimensionFolding(dace_transformation.SingleStateTransformation): fusion is rejected. Folding the value into the Memlets removes that blind spot. Note: - The dimension itself is kept, in contrast to DaCe's `TrivialMapElimination` - which removes it. Removing it changes the shape of the Map and prevents it - from being scheduled on the GPU, whereas folding only rewrites the uses of - the parameter and leaves the iteration space, the schedule and any blocking - applied later untouched. + 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. @@ -61,7 +58,7 @@ def _single_iteration_parameters(map_: dace_nodes.Map) -> dict[str, Any]: return { param: rng[0] for param, rng in zip(map_.params, map_.range.ranges) - if (rng[0] == rng[1]) == True and (rng[2] == 1) == True # noqa: E712 [true-false-comparison] # SymPy comparison + if (rng[0] == rng[1]) == True # noqa: E712 [true-false-comparison] # SymPy comparison } def can_be_applied( @@ -99,11 +96,10 @@ def apply(self, graph: Union[dace.SDFGState, dace.SDFG], sdfg: dace.SDFG) -> Non replacements = self._single_iteration_parameters(map_entry.map) scope = graph.scope_subgraph(map_entry, include_entry=True, include_exit=True) - # `replace()` would also rewrite the Map's own parameters and range, which + # `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) - for param, value in replacements.items(): - scope.replace(param, value) + 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 index 096ebfbf5e..d184c27aa6 100644 --- 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 @@ -52,6 +52,7 @@ def test_trivial_map_dimension_folding(): 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, @@ -69,27 +70,18 @@ def test_trivial_map_dimension_folding(): assert map_entry.map.params == ["__i", "__j"] assert str(map_entry.map.range[1][0]) == "7" assert all( - "__j" not in str(edge.data.subset) + "__j" not in str(subset) for state in sdfg.states() for edge in state.edges() - if edge.data is not None and edge.data.subset is not None + 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_is_a_fixpoint(): - """A second application must not be possible, else the pass would not terminate.""" - sdfg = _mk_single_iteration_sdfg() - assert ( - sdfg.apply_transformations_repeated( - gtx_transformations.TrivialMapDimensionFolding(), validate=True - ) - == 1 - ) - - 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")) From 586c5b4d0854ee0c81f8380c1c2899764edad88e Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Mon, 17 Aug 2026 11:12:04 +0200 Subject: [PATCH 3/3] refactor[next-dace]: address review comments on 2799 - point the property description at the class docstring - drop the note on `include_entry`, it holds for `include_exit` just as much Co-Authored-By: Claude Opus 5 --- .../dace/transformations/trivial_map_dimension_folding.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) 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 index 111d243aec..9c3542e6d9 100644 --- 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 @@ -40,7 +40,7 @@ class TrivialMapDimensionFolding(dace_transformation.SingleStateTransformation): only_toplevel_maps = dace_properties.Property( dtype=bool, default=False, - desc="Only process Maps that are on the top level.", + desc="See docs.", ) def __init__(self, only_toplevel_maps: Optional[bool] = None, **kwargs: Any) -> None: @@ -78,8 +78,6 @@ def can_be_applied( # Only apply if a parameter is still referenced, otherwise the transformation # would apply again and again on the same Map. - # NOTE: `include_entry` is needed because the uses are on the out edges of - # the MapEntry, which are not part of the scope subgraph otherwise. 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(