From eca948e3be46ff2fdd7228052834b151c80fb8d4 Mon Sep 17 00:00:00 2001 From: Till Ehrengruber Date: Thu, 1 Oct 2026 14:34:50 +0000 Subject: [PATCH] refactor[next]: parameterize element enumeration in 'tree_map' MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Traversing type trees with `tree_map` only worked because `ts.TupleType` and `ts.NamedCollectionType` implemented `__iter__`/`__len__` — container protocols bolted onto plain data classes solely to serve the traversal. Make the enumeration of a node's children a parameter of the traversal instead, so the dunders can be removed. `tree_map` now takes two keyword-only parameters: - `collection_elements` (default `iter`) decomposes a collection into its elements, the dual of `result_collection_constructor`. - `collection_keys` (default an unbounded `itertools.count()`) supplies the path components for `with_path_arg`, so elements can be addressed by something other than their position. Together with `collection_type` these are the three facets of a traversable collection: membership, decomposition and reconstruction. Call sites that spelled out `collection_type=ts....` now go through `type_info.tree_map_type`, which is additionally usable as a (parametrized) decorator so that the older sites can adopt it. Removing `__len__` also means an empty `TupleType` is no longer falsy, which fixes the `if node.type:` / `assert node.type` checks in type inference that were meant as "is the type set". --- src/gt4py/next/ffront/past_to_itir.py | 3 +- src/gt4py/next/ffront/type_info.py | 3 +- src/gt4py/next/field_utils.py | 7 +- .../iterator/transforms/expand_tuple_maps.py | 8 +- .../next/iterator/transforms/infer_domain.py | 6 +- .../iterator/type_system/type_synthesizer.py | 36 +++--- src/gt4py/next/type_system/type_info.py | 25 +++- .../next/type_system/type_specifications.py | 15 +-- src/gt4py/next/utils.py | 67 ++++++++++- tests/next_tests/unit_tests/test_utils.py | 108 ++++++++++++++++++ 10 files changed, 220 insertions(+), 58 deletions(-) diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 3febb910ef..bc4648dfe2 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -457,8 +457,7 @@ def _visit_stencil_call_out_arg( "Unexpected 'out' argument. Must be a 'past.Subscript', 'past.Name' or 'past.TupleExpr' node." ) - @utils.tree_map( - collection_type=ts.COLLECTION_TYPE_SPECS, + @type_info.tree_map_type( with_path_arg=True, unpack=True, result_collection_constructor=lambda _, elts: im.make_tuple(*elts), diff --git a/src/gt4py/next/ffront/type_info.py b/src/gt4py/next/ffront/type_info.py index cfe1a51abb..e7280d51b2 100644 --- a/src/gt4py/next/ffront/type_info.py +++ b/src/gt4py/next/ffront/type_info.py @@ -20,9 +20,8 @@ named_collections_to_tuple_types = cast( Callable[..., ts.TupleType], - utils.tree_map( + type_info.tree_map_type( lambda x: x, - collection_type=ts.COLLECTION_TYPE_SPECS, result_collection_constructor=lambda _, elems: ts.TupleType(types=list(elems)), ), ) diff --git a/src/gt4py/next/field_utils.py b/src/gt4py/next/field_utils.py index 460d4caff0..97e14ae0ef 100644 --- a/src/gt4py/next/field_utils.py +++ b/src/gt4py/next/field_utils.py @@ -13,7 +13,7 @@ from gt4py._core import definitions as core_defs from gt4py.eve.extended_typing import NestedTuple from gt4py.next import common, named_collections, utils -from gt4py.next.type_system import type_specifications as ts, type_translation +from gt4py.next.type_system import type_info, type_specifications as ts, type_translation try: @@ -61,10 +61,7 @@ def _constructor( return named_collections.make_named_collection_constructor_from_type_spec(type_)(elems) return tuple(elems) - @utils.tree_map( - collection_type=ts.COLLECTION_TYPE_SPECS, - result_collection_constructor=_constructor, - ) + @type_info.tree_map_type(result_collection_constructor=_constructor) def impl(type_: ts.ScalarType) -> common.MutableField: res = common._field( xp.empty(domain.shape, dtype=xp.dtype(type_translation.as_dtype(type_).scalar_type)), diff --git a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py index 6b705a29dc..56968f77de 100644 --- a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py +++ b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py @@ -14,16 +14,14 @@ from gt4py.next.iterator import ir as itir from gt4py.next.iterator.ir_utils import common_pattern_matcher as cpm, ir_makers as im from gt4py.next.iterator.type_system import inference as itir_inference -from gt4py.next.type_system import type_specifications as ts +from gt4py.next.type_system import type_info, type_specifications as ts def _tree_map_tuple_body(f: itir.Expr, tup_expr: itir.Expr, tup_type: ts.TupleType) -> itir.Expr: """Recursively expand `tree_map_tuple(f)(t)` into `make_tuple` calls.""" - @utils.tree_map( - collection_type=ts.TupleType, - result_collection_constructor=lambda _, elts: im.make_tuple(*elts), - with_path_arg=True, + @type_info.tree_map_type( + result_collection_constructor=lambda _, elts: im.make_tuple(*elts), with_path_arg=True ) def mapper(_el_type, path): return im.call(f)(functools.reduce(lambda expr, i: im.tuple_get(i, expr), path, tup_expr)) diff --git a/src/gt4py/next/iterator/transforms/infer_domain.py b/src/gt4py/next/iterator/transforms/infer_domain.py index c7e7d4f448..da1bece83b 100644 --- a/src/gt4py/next/iterator/transforms/infer_domain.py +++ b/src/gt4py/next/iterator/transforms/infer_domain.py @@ -475,9 +475,9 @@ def infer_expr( expr, offset_provider_type=common.offset_provider_to_type(offset_provider) ) el_types, domain = gtx_utils.equalize_tuple_structure( - gtx_utils.tree_map( - collection_type=ts.TupleType, result_collection_constructor=lambda _, elts: tuple(elts) - )(lambda x: x)(expr.type), + type_info.tree_map_type( + lambda x: x, result_collection_constructor=lambda _, elts: tuple(elts) + )(expr.type), domain, fill_value=DomainAccessDescriptor.NEVER, # el_types already has the right structure, we only want to change domain diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index a3c98a96fb..453f8803e1 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -15,7 +15,7 @@ from gt4py.eve import utils as eve_utils from gt4py.eve.extended_typing import Callable, Iterable, Optional, Union -from gt4py.next import common, utils +from gt4py.next import common from gt4py.next.iterator import builtins, ir as itir from gt4py.next.iterator.ir_utils import misc as ir_misc from gt4py.next.iterator.type_system import type_specifications as it_ts @@ -202,10 +202,7 @@ def if_( pred: ts.ScalarType | ts.DeferredType, true_branch: ts.DataType, false_branch: ts.DataType ) -> ts.DataType: if isinstance(true_branch, ts.TupleType) and isinstance(false_branch, ts.TupleType): - return utils.tree_map( - collection_type=ts.TupleType, - result_collection_constructor=lambda _, elts: ts.TupleType(types=[*elts]), - )(functools.partial(if_, pred))(true_branch, false_branch) + return type_info.tree_map_type(functools.partial(if_, pred))(true_branch, false_branch) assert not isinstance(true_branch, ts.TupleType) and not isinstance(false_branch, ts.TupleType) assert isinstance(pred, ts.DeferredType) or ( @@ -271,10 +268,7 @@ def concat_where( if isinstance(true_field, ts.DeferredType) or isinstance(false_field, ts.DeferredType): return ts.DeferredType(constraint=None) - @utils.tree_map( - collection_type=ts.TupleType, - result_collection_constructor=lambda _, elts: ts.TupleType(types=list(elts)), - ) + @type_info.tree_map_type def deduce_return_type(tb: ts.FieldType | ts.ScalarType, fb: ts.FieldType | ts.ScalarType): if any(isinstance(b, ts.DeferredType) for b in [tb, fb]): return ts.DeferredType(constraint=ts.FieldType) @@ -292,17 +286,20 @@ def deduce_return_type(tb: ts.FieldType | ts.ScalarType, fb: ts.FieldType | ts.S return_type = ts.FieldType(dims=return_dims, dtype=dtype) return return_type - return deduce_return_type(true_field, false_field) + result = deduce_return_type(true_field, false_field) + assert isinstance(result, (ts.FieldType, ts.TupleType, ts.DeferredType)) + return result @_register_builtin_type_synthesizer def broadcast( - arg: ts.FieldType | ts.ScalarType | ts.DeferredType, dims: tuple[ts.DimensionType] + arg: ts.FieldType | ts.ScalarType | ts.DeferredType, dims: ts.TupleType ) -> ts.FieldType | ts.DeferredType: if isinstance(arg, ts.DeferredType): return arg - dims_ = [dim.dim for dim in dims] + assert all(isinstance(dim, ts.DimensionType) for dim in dims.types) + dims_ = [cast(ts.DimensionType, dim).dim for dim in dims.types] if isinstance(arg, ts.FieldType): dtype = arg.dtype @@ -422,13 +419,14 @@ def _canonicalize_nb_fields( """ match input_: case tuple() | ts.TupleType(): + elements = input_.types if isinstance(input_, ts.TupleType) else input_ assert all( - isinstance(field, (ts.ScalarType, ts.FieldType, ts.TupleType)) for field in input_ + isinstance(field, (ts.ScalarType, ts.FieldType, ts.TupleType)) for field in elements ) return ts.TupleType( types=[ _canonicalize_nb_fields(cast(ts.FieldType | ts.TupleType, field)) - for field in input_ + for field in elements ] ) case ts.FieldType(): @@ -587,7 +585,7 @@ def applied_as_fieldop( output_dims: list[common.Dimension] = [] if offset_provider_type is not None and shift_sequences_per_param is not None: for field, shift_sequences in zip( - new_fields, shift_sequences_per_param, strict=True + new_fields.types, shift_sequences_per_param, strict=True ): for el in type_info.primitive_constituents(field): input_dims = type_info.extract_dims(el) @@ -608,7 +606,7 @@ def applied_as_fieldop( return ts.DeferredType(constraint=None) stencil_return = stencil( - *(_convert_as_fieldop_input_to_iterator(domain, field) for field in new_fields), + *(_convert_as_fieldop_input_to_iterator(domain, field) for field in new_fields.types), offset_provider_type=offset_provider_type, ) @@ -685,11 +683,7 @@ def applied_map( bound_op = functools.partial(op, offset_provider_type=offset_provider_type) if recursive: - return utils.tree_map( # type: ignore[return-value] - bound_op, - collection_type=ts.TupleType, - result_collection_constructor=lambda _, elts: ts.TupleType(types=[*elts]), - )(arg) + return type_info.tree_map_type(bound_op)(arg) # type: ignore[return-value] # Non-recursive: apply `op` once per top-level element. return ts.TupleType(types=[bound_op(el) for el in arg.types]) diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index ac8467a5f6..8ac49f874e 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -162,16 +162,39 @@ def tree_map_type( ) -> Callable[..., _T | _C]: ... +@overload def tree_map_type( - fun: Callable[..., _T], + *, + result_collection_constructor: Callable[..., Any] = ..., + with_path_arg: bool = ..., + unpack: bool = ..., +) -> Callable[[Callable[..., Any]], Callable[..., Any]]: ... + + +def tree_map_type( + fun: Callable[..., _T] | None = None, *, result_collection_constructor: Callable[..., Any] = tree_map_type_constructor, with_path_arg: bool = False, unpack: bool = False, ) -> Callable[..., Any]: + """ + `tree_map` specialized for collection type specs, see `gt4py.next.utils.tree_map`. + + Can be used directly or as a (possibly parametrized) decorator. By default the result + collection is of the same kind as the traversed collection type spec. + """ + if fun is None: + return functools.partial( + tree_map_type, + result_collection_constructor=result_collection_constructor, + with_path_arg=with_path_arg, + unpack=unpack, + ) return next_utils.tree_map( fun, collection_type=ts.COLLECTION_TYPE_SPECS, + collection_elements=lambda t: t.types, result_collection_constructor=result_collection_constructor, with_path_arg=with_path_arg, unpack=unpack, diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 59ac40f0f3..8016810d15 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -8,7 +8,7 @@ from __future__ import annotations -from typing import Final, Iterator, Optional, Sequence, TypeVar +from typing import Final, Optional, Sequence, TypeVar from gt4py.eve import ( datamodels as eve_datamodels, @@ -141,12 +141,6 @@ class TupleType(DataType): def __str__(self) -> str: return f"tuple[{', '.join(map(str, self.types))}]" - def __iter__(self) -> Iterator[DataType | DimensionType | DeferredType]: - yield from self.types - - def __len__(self) -> int: - return len(self.types) - class AnyPythonType: """Marker type representing any Python type which cannot be used for instantiation. @@ -190,13 +184,6 @@ def __getattr__(self, name: str) -> DataType | DimensionType | DeferredType: def __str__(self) -> str: return f"NamedTuple{{{', '.join(f'{k}: {v}' for k, v in zip(self.keys, self.types))}}}" - def __iter__(self) -> Iterator[DataType | DimensionType | DeferredType]: - # Note: Unlike `Mapping`s, we iterate the values (not the keys) by default. - yield from self.types - - def __len__(self) -> int: - return len(self.types) - CollectionTypeSpecT = TypeVar("CollectionTypeSpecT", TupleType, NamedCollectionType) CollectionTypeSpec = TupleType | NamedCollectionType diff --git a/src/gt4py/next/utils.py b/src/gt4py/next/utils.py index 6b15f7b694..54491d3663 100644 --- a/src/gt4py/next/utils.py +++ b/src/gt4py/next/utils.py @@ -12,7 +12,7 @@ import inspect import itertools import types -from collections.abc import Callable +from collections.abc import Callable, Iterable from typing import ( Any, ClassVar, @@ -104,6 +104,8 @@ def tree_map( fun: Callable[_P, _R], *, collection_type: type | tuple[type, ...] = tuple, + collection_elements: Callable[[Any], Iterable[Any]] = iter, + collection_keys: Callable[[Any], Iterable[Any]] = lambda _: itertools.count(), result_collection_constructor: Optional[Callable] = None, unpack: bool = False, with_path_arg: bool = False, @@ -114,6 +116,8 @@ def tree_map( def tree_map( *, collection_type: type | tuple[type, ...] = tuple, + collection_elements: Callable[[Any], Iterable[Any]] = iter, + collection_keys: Callable[[Any], Iterable[Any]] = lambda _: itertools.count(), result_collection_constructor: Optional[Callable] = None, unpack: bool = False, with_path_arg: bool = False, @@ -126,6 +130,8 @@ def tree_map( fun: Optional[Callable[_P, _R]] = None, *, collection_type: type | tuple[type, ...] = tuple, + collection_elements: Callable[[Any], Iterable[Any]] = iter, + collection_keys: Callable[[Any], Iterable[Any]] = lambda _: itertools.count(), result_collection_constructor: Optional[Callable] = None, unpack: bool = False, with_path_arg: bool = False, @@ -136,6 +142,16 @@ def tree_map( Args: fun: Function to apply to each entry of the collection. collection_type: Type of the collection to be traversed. Can be a single type or a tuple of types. + collection_elements: Decompose a collection into its elements. Dual to + `result_collection_constructor`: the three facets of a traversable collection are + membership (`collection_type`), decomposition (`collection_elements`) and + reconstruction (`result_collection_constructor`). Defaults to `iter`, i.e. the + collections are assumed to implement the iterator protocol. + collection_keys: Keys of a collection's elements, used as the path components when + `with_path_arg` is set. May be an unbounded iterator, as the default is, which + addresses elements by their position. Note that the keys are not needed to + reconstruct a collection, since `result_collection_constructor` receives the original + collection and can recover them from there. result_collection_constructor: Type of the collection to be returned. If `None` the same type as `collection_type` is used. unpack: Replicate tuple structure returned from `fun` to the mapped result, i.e. return tuple of result collections instead of result collections of tuples. @@ -184,6 +200,36 @@ def tree_map( ((2, 3), 4) >>> squared ((4, 9), 16) + + Collections that do not implement the iterator protocol are traversed by passing a + custom `collection_elements` decomposition: + + >>> import dataclasses + >>> @dataclasses.dataclass + ... class Node: + ... children: dict + >>> tree_map( + ... collection_type=Node, + ... collection_elements=lambda node: node.children.values(), + ... result_collection_constructor=lambda value, elts: Node( + ... children=dict(zip(value.children.keys(), elts)) + ... ), + ... )(lambda x: x + 1)(Node({"a": Node({"b": 1})})) + Node(children={'a': Node(children={'b': 2})}) + + Elements can be addressed by something other than their position by additionally + passing `collection_keys`: + + >>> tree_map( + ... collection_type=Node, + ... collection_elements=lambda node: node.children.values(), + ... collection_keys=lambda node: node.children.keys(), + ... result_collection_constructor=lambda value, elts: dict( + ... zip(value.children.keys(), elts) + ... ), + ... with_path_arg=True, + ... )(lambda x, path: path)(Node({"a": Node({"b": 1})})) + {'a': {'b': ('a', 'b')}} """ if result_collection_constructor is None: @@ -199,20 +245,29 @@ def tree_map( @functools.wraps(fun) def impl(*args: Any | tuple[Any | tuple, ...]) -> _R | tuple[_R | tuple, ...]: if isinstance(args[0], collection_type): - first_arg: Any = args[0] non_path_args: Sequence[Any] + path: tuple[Any, ...] = () if with_path_arg: *non_path_args, path = args - args = (*non_path_args, tuple((*path, i) for i in range(len(first_arg)))) else: non_path_args = args assert all(isinstance(arg, collection_type) for arg in non_path_args) - assert all(len(first_arg) == len(arg) for arg in non_path_args) + first_arg: Any = non_path_args[0] + elements_per_arg = [tuple(collection_elements(arg)) for arg in non_path_args] + + zipped_args: list[Iterable[Any]] = [*elements_per_arg] + if with_path_arg: + # Note: `collection_keys` may be an unbounded iterator (as the default is), + # hence `zip` below stops at the shortest argument and the structures are + # only compared after the mapping. + zipped_args.append((*path, key) for key in collection_keys(first_arg)) + assert result_collection_constructor is not None ctor = functools.partial(result_collection_constructor, first_arg) - mapped = [impl(*arg) for arg in zip(*args)] + mapped = [impl(*elements) for elements in zip(*zipped_args)] + assert all(len(mapped) == len(elements) for elements in elements_per_arg) if unpack: return tuple(map(ctor, zip(*mapped))) else: @@ -230,6 +285,8 @@ def impl(*args: Any | tuple[Any | tuple, ...]) -> _R | tuple[_R | tuple, ...]: return functools.partial( tree_map, collection_type=collection_type, + collection_elements=collection_elements, + collection_keys=collection_keys, result_collection_constructor=result_collection_constructor, unpack=unpack, with_path_arg=with_path_arg, diff --git a/tests/next_tests/unit_tests/test_utils.py b/tests/next_tests/unit_tests/test_utils.py index 6068d06f3d..02c48fd725 100644 --- a/tests/next_tests/unit_tests/test_utils.py +++ b/tests/next_tests/unit_tests/test_utils.py @@ -652,6 +652,114 @@ def testee(x): assert testee(((1, 2), 3)) == [[2, 3], 4] +@dataclasses.dataclass +class _Node: + """Collection whose children are only reachable through an attribute, not by iteration.""" + + children: list + + +def _node_elements(node: _Node) -> list: + return node.children + + +def _node_constructor(_, elements) -> _Node: + return _Node(children=list(elements)) + + +@dataclasses.dataclass +class _NamedNode: + """Collection whose children are addressed by name instead of by position.""" + + children: dict + + +def _named_node_constructor(value: _NamedNode, elements) -> _NamedNode: + return _NamedNode(children=dict(zip(value.children.keys(), elements, strict=True))) + + +def test_tree_map_custom_collection_elements(): + @utils.tree_map( + collection_type=_Node, + collection_elements=_node_elements, + result_collection_constructor=_node_constructor, + ) + def testee(x): + return x + 1 + + assert testee(_Node([_Node([1, 2]), 3])) == _Node([_Node([2, 3]), 4]) + + +def test_tree_map_custom_collection_elements_multi_arg(): + @utils.tree_map( + collection_type=_Node, + collection_elements=_node_elements, + result_collection_constructor=_node_constructor, + ) + def testee(x, y): + return x + y + + assert testee(_Node([_Node([1, 2]), 3]), _Node([_Node([4, 5]), 6])) == _Node([_Node([5, 7]), 9]) + + +def test_tree_map_custom_collection_elements_with_path_arg(): + @utils.tree_map( + collection_type=_Node, + collection_elements=_node_elements, + result_collection_constructor=_node_constructor, + with_path_arg=True, + ) + def testee(x, path): + return (x, path) + + assert testee(_Node([_Node([1, 2]), 3])) == _Node( + [_Node([(1, (0, 0)), (2, (0, 1))]), (3, (1,))] + ) + + +def test_tree_map_custom_collection_elements_unpack(): + @utils.tree_map( + collection_type=_Node, + collection_elements=_node_elements, + result_collection_constructor=_node_constructor, + unpack=True, + ) + def testee(x): + return (x, x**2) + + assert testee(_Node([_Node([2, 3]), 4])) == ( + _Node([_Node([2, 3]), 4]), + _Node([_Node([4, 9]), 16]), + ) + + +def test_tree_map_custom_collection_keys(): + """`collection_keys` makes up the path, independently of the element decomposition.""" + + @utils.tree_map( + collection_type=_NamedNode, + collection_elements=lambda node: node.children.values(), + collection_keys=lambda node: node.children.keys(), + result_collection_constructor=_named_node_constructor, + with_path_arg=True, + ) + def testee(x, path): + return path + + assert testee(_NamedNode({"a": _NamedNode({"b": 1}), "c": 2})) == _NamedNode( + {"a": _NamedNode({"b": ("a", "b")}), "c": ("c",)} + ) + + +def test_tree_map_structure_mismatch(): + @utils.tree_map(with_path_arg=True) + def testee(x, y, path): + return (x + y, path) + + with pytest.raises(AssertionError): + testee((1, 2), (3, 4, 5)) + + def test_tree_map_multiple_input_types(): @utils.tree_map( collection_type=(list, tuple),