Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -123,6 +124,7 @@
"SplitAccessNode",
"SplitConsumerMemlet",
"TransientMemoryMode",
"TrivialMapDimensionFolding",
"VerticalMapFusionCallback",
"VerticalMapSplitCallback",
"constants",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
@@ -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
Original file line number Diff line number Diff line change
@@ -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
)