Skip to content
Open
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
3 changes: 1 addition & 2 deletions src/gt4py/next/ffront/past_to_itir.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
3 changes: 1 addition & 2 deletions src/gt4py/next/ffront/type_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)),
),
)
Expand Down
7 changes: 2 additions & 5 deletions src/gt4py/next/field_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
from gt4py._core import definitions as core_defs
from gt4py.eve.xtyping 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:
Expand Down Expand Up @@ -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)),
Expand Down
8 changes: 3 additions & 5 deletions src/gt4py/next/iterator/transforms/expand_tuple_maps.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
6 changes: 3 additions & 3 deletions src/gt4py/next/iterator/transforms/infer_domain.py
Original file line number Diff line number Diff line change
Expand Up @@ -490,9 +490,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
Expand Down
36 changes: 15 additions & 21 deletions src/gt4py/next/iterator/type_system/type_synthesizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
from typing import Optional, TypeVar, Union, cast, overload

from gt4py.eve import utils as eve_utils
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
Expand Down Expand Up @@ -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 (
Expand Down Expand Up @@ -273,10 +270,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)
Expand All @@ -287,17 +281,20 @@ def deduce_return_type(tb: ts.FieldType | ts.ScalarType, fb: ts.FieldType | ts.S
dtype=type_info.extract_dtype(promoted),
)

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
Expand Down Expand Up @@ -417,13 +414,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():
Expand Down Expand Up @@ -582,7 +580,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)
Expand All @@ -603,7 +601,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,
)

Expand Down Expand Up @@ -680,11 +678,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])
Expand Down
25 changes: 24 additions & 1 deletion src/gt4py/next/type_system/type_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,16 +163,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,
Expand Down
15 changes: 1 addition & 14 deletions src/gt4py/next/type_system/type_specifications.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
from __future__ import annotations

import typing
from typing import Final, Iterator, Optional, Sequence, TypeVar
from typing import Final, Optional, Sequence, TypeVar

from gt4py.eve import datamodels as eve_datamodels, type_definitions as eve_types
from gt4py.next import common
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
67 changes: 62 additions & 5 deletions src/gt4py/next/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -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.
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand All @@ -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,
Expand Down
Loading
Loading