From 38e83a55a41aa01eaa257e03bf7cb94374da1ee1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 23:39:56 +0200 Subject: [PATCH 01/13] feat[next]: NeighborConnectivity declarations and local dimensions that know their owner --- .../ADRs/next/0029-Connectivities_As_Types.md | 122 +++++++ docs/development/ADRs/next/README.md | 3 +- src/gt4py/next/__init__.py | 4 + src/gt4py/next/common.py | 283 +++++++++++++++- src/gt4py/next/embedded/nd_array_field.py | 15 +- .../ffront/foast_passes/type_deduction.py | 12 +- src/gt4py/next/ffront/past_to_itir.py | 2 +- src/gt4py/next/ffront/transform_utils.py | 8 +- src/gt4py/next/iterator/embedded.py | 25 +- .../test_neighbor_connectivity.py | 104 ++++++ .../unit_tests/test_neighbor_connectivity.py | 314 ++++++++++++++++++ typing_tests/test_next.yaml | 69 ++++ 12 files changed, 937 insertions(+), 24 deletions(-) create mode 100644 docs/development/ADRs/next/0029-Connectivities_As_Types.md create mode 100644 tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py create mode 100644 tests/next_tests/unit_tests/test_neighbor_connectivity.py diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md new file mode 100644 index 0000000000..2fdcc101e8 --- /dev/null +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -0,0 +1,122 @@ +--- +tags: [] +--- + +# Connectivities as Types + +- **Status**: proposed +- **Authors**: Enrique González Paredes (@egparedes) +- **Created**: 2026-09-21 +- **Updated**: 2026-09-21 + +A neighbor connectivity is declared as a **class**, and its local dimension as a +class **nested** in it: + +```python +class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): + class Local(gtx.LocalDimensionIndex): ... + + +@gtx.field_operator +def f(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return neighbor_sum(a(V2E), axis=V2E.Local) + a(V2E[0]) +``` + +The declaration is written in DSL code, owns its local dimension, and states the +constraints the neighbor table bound at call time has to satisfy. It holds no +data. It builds on [ADR 0028](0028-Dimensions_As_Nominal_Types.md): the +connectivity, like a dimension, is identified by its type, and `V2E.Local` is an +ordinary dimension class with the tag `.V2E.Local`. + +## Context + +An unstructured connectivity used to be spelled by four independently authored +names that had to agree, none of them checked against the others: the +`FieldOffset` tag, the Python variable it was bound to, the local dimension's +name and the offset-provider key. The `V2EDim = Dimension("V2E")` convention made +all four equal, which hid which one each execution path actually used; the +regression tests in `test_offset_dimensions_names.py` break the convention one +name at a time. Nothing tied a local dimension to the table it indexes, so the +backends recovered that link by string equality, and the table's shape, codomain +and skip values were never checked against the `FieldOffset` declaration. + +## Decision + +### The declaration + +- `NeighborConnectivity[Origin, Codomain]` is a PEP 695 generic whose subclasses + are declarations: for each `Origin` element, a list of `Codomain` neighbors. Its + metaclass, `ConnectivityMeta`, forbids instantiation. +- The local dimension is the nested class `Local`, a subclass of + `LocalDimensionIndex`. Declaring it is required, and `NeighborConnectivity` + sets `Local.owner` to the connectivity when the class is created. A local + dimension can have at most one owner. +- A local dimension with no table, such as the coefficient axis of a fixed-size + stencil, is declared on its own: `class LsqCoeff(LocalDimensionIndex, size=3)`. + Its `owner` is `None`. +- `max_neighbors` and `min_neighbors` are optional class keywords, not type + parameters: Python has no integer type parameters, and nothing static needs + the count. A declared count is a constraint on the bound table; an undeclared + one is taken from the table. `min_neighbors < max_neighbors` means that the + table must use skip values. +- `common.check_neighbor_table(V2E, table)` checks a table, or just its type + (which is all an ahead-of-time compilation has), against the declaration: + the domain is `(Origin, V2E.Local)`, the codomain is `Codomain`, the dtype is + integral, and the neighbor counts and skip values agree. + +### `NeighborConnectivity` is not a `Connectivity` + +`common.Connectivity` is a *data* protocol (`ndarray`, `domain`, `asnumpy`); a +declaration holds no data. The neighbor table stays a `Connectivity` +implementation, and the declaration is only the type the table is checked +against. `Field.premap` and `Field.__call__` accept either, as they already +accepted a `FieldOffset`, which is not a `Connectivity` either. + +### `LocalDimensionIndex` subclasses `DimensionIndex` + +A separate root would force every `type[DimensionIndex]` annotation in the tree +(`ts.FieldType.dims`, `Domain`, `ConnectivityType.domain`, ...) to widen, and +would then accept local dimensions wherever a primary one is meant anyway. The +tree already distinguishes local dimensions by a runtime `kind` check, so it +keeps doing so; generic constructors whose parameter must be a primary dimension +(`NeighborConnectivity[Origin, Codomain]`, `Staggered[D]`) check it at runtime. + +### How the base reaches `Local` + +The base class declares `Local: ClassVar[type[LocalDimensionIndex]]` as an +annotation only. A real nested class on the base would be an incompatible +override for pyright in every declaration. The annotation makes `conn.Local` a +value of type `type[LocalDimensionIndex]` for library code taking any +connectivity; it does not make `conn.Local` usable as an *annotation* when +`conn` is generic, which no checker allows. Code that needs to name a local +dimension generically uses a `TypeVar` bound to `LocalDimensionIndex`. + +### Frontend integration + +A declaration is typed like the `FieldOffset` it replaces: `V2E.__gt_type__()` +is the `ts.OffsetType` of the derived offset `(Codomain -> (Origin, Local))`, +whose tag is **the local dimension's tag**, `V2E.Local.tag`. This is the single +string that shifts, neighbor reductions and sparse arguments already use to find +the table in the offset provider, so existing backends need no change. +`V2E.Local` inside DSL code types as that local dimension. `V2E[i]` subscripts the +metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. + +## Consequences + +- An unstructured connectivity is spelled once. The provider key, the offset tag + and the local dimension are all derived from the declaration. +- A table bound to a connectivity can be checked against its declaration. +- The frontend needs one special case: `V2E.Local` is resolved from the offset + type, because the type of `V2E` is not the class. +- `FieldOffset` remains during migration; a `FieldOffset` and a + `NeighborConnectivity` sharing a local dimension are interchangeable. + +## Alternatives considered + +- **The local dimension generated by the metaclass**, e.g. `V2E.Local` created + from `V2E`'s name. Type checkers cannot see a generated class, so it could not + be used in `Field[Dims[Vertex, V2E.Local], ...]`. +- **Neighbor counts as type parameters.** Python has no integer type parameters, + and a `Literal[6]` argument would add a type parameter nothing statically uses. +- **`NeighborConnectivity` as a `Connectivity` subclass.** Mixes the + declaration with the data protocol; see above. diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index 4f1661d745..4946d31972 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -22,7 +22,8 @@ Writing a new ADR is simple: - [0021 - Argument Descriptors](0021-Argument-Descriptors.md) - [0023 - Fingerprinting](0023-Fingerprinting.md) - [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md) -- [0029 - Dimensions as Nominal Types](0029-Dimensions_As_Nominal_Types.md) +- [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md) +- [0029 - Connectivities as Types](0029-Connectivities_As_Types.md) ### Frontend and Parsing #frontend diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index b8a7bf5143..c41ce38680 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -33,6 +33,8 @@ Domain, Field, GridType, + LocalDimensionIndex, + NeighborConnectivity, Staggered, UnitRange, as_non_staggered, @@ -124,6 +126,8 @@ "AnyCartesianAxisIndex", "CartesianAxisIndex", "DimensionKind", + "LocalDimensionIndex", + "NeighborConnectivity", "Staggered", "resolve", "Dims", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index a5486c024f..dc4b9eda7c 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -16,6 +16,7 @@ import functools import importlib import math +import numbers import re import sys import types @@ -1107,7 +1108,9 @@ def asnumpy(self) -> np.ndarray: ... def as_scalar(self) -> core_defs.ScalarT: ... @abc.abstractmethod - def premap(self, index_field: Connectivity | fbuiltins.FieldOffset) -> Field: ... + def premap( + self, index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity] + ) -> Field: ... @abc.abstractmethod def restrict(self, item: AnyIndexSpec) -> Self: ... @@ -1115,8 +1118,8 @@ def restrict(self, item: AnyIndexSpec) -> Self: ... @abc.abstractmethod def __call__( self, - index_field: Connectivity | fbuiltins.FieldOffset, - *args: Connectivity | fbuiltins.FieldOffset, + index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], ) -> Field: ... @abc.abstractmethod @@ -1637,8 +1640,8 @@ def inverse_image(self, image_range: UnitRange | NamedRange) -> Sequence[NamedRa def premap( self, - index_field: Connectivity | fbuiltins.FieldOffset, - *args: Connectivity | fbuiltins.FieldOffset, + index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], ) -> Connectivity: raise NotImplementedError() @@ -1811,6 +1814,8 @@ def __getitem__(cls, base: Dimension) -> Dimension: ) if not isinstance(base, DimensionMeta): raise TypeError(f"'Staggered' expects a dimension, got '{base!r}'.") + if base.kind is DimensionKind.LOCAL: + raise TypeError(f"'{base.__qualname__}' is a local dimension and cannot be staggered.") if is_staggered(base): raise TypeError( f"'{base.__qualname__}' is already staggered; a dimension cannot be staggered twice." @@ -1961,3 +1966,271 @@ def connectivity_for_cartesian_shift(dim: Dimension, offset: int | float) -> Car else: assert staggered_offset == 0 return CartesianConnectivity(dim, int(integral_offset), codomain=dim) + + +class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): + """ + A local dimension: the axis that runs over the neighbors of one element. + + A local dimension is declared either inside a `NeighborConnectivity`, as its nested `Local` + class, or on its own for a local axis that indexes no table (`owner is None`), such as the + coefficients of a fixed-size stencil: + + >>> class LsqCoeff(LocalDimensionIndex, size=3): ... + >>> LsqCoeff.kind, LsqCoeff.owner, LsqCoeff.max_neighbors + (, None, 3) + + Neighbor counts are optional. A declared count is a constraint the bound table has to + satisfy (see `check_neighbor_table`); an undeclared one is taken from the table. + """ + + __slots__ = () + + #: The connectivity this dimension is the local axis of, or `None` if it indexes no table. + #: Set by `NeighborConnectivity` when the connectivity is declared. + owner: ClassVar[Optional[type[NeighborConnectivity]]] = None + #: Number of entries per element, i.e. the table's second extent, if declared. + max_neighbors: ClassVar[Optional[int]] = None + #: Least number of *valid* neighbors of any element, if declared. Fewer than + #: `max_neighbors` means the table pads with skip values. + min_neighbors: ClassVar[Optional[int]] = None + + def __init_subclass__( + cls, + /, + *, + size: Optional[int] = None, + kind: Optional[DimensionKind] = None, + **kwargs: Any, + ) -> None: + if kind is not None and kind is not DimensionKind.LOCAL: + raise TypeError( + f"'{cls.__qualname__}' is a local dimension and cannot have kind '{kind}'." + ) + super().__init_subclass__(**kwargs) + # NOTE: reset rather than inherited: a subclass of an owned local dimension is a + # different dimension, and does not index its parent's table. + cls.owner = None + cls.max_neighbors = cls.min_neighbors = _check_neighbor_count(cls, "size", size) + + +def _check_neighbor_count(owner: type, name: str, count: Optional[int]) -> Optional[int]: + if count is None: + return None + if not isinstance(count, numbers.Integral) or isinstance(count, bool) or count < 0: + raise TypeError( + f"'{owner.__qualname__}': '{name}' must be a non-negative integer, got '{count!r}'." + ) + return int(count) + + +class ConnectivityMeta(type): + """ + Metaclass of `NeighborConnectivity` declarations. + + A connectivity declaration is a class that is never instantiated. It is written in DSL code + (`a(V2E)`, `a(V2E[0])`), and it is what the neighbor table bound at call time must match. + """ + + Local: type[LocalDimensionIndex] + origin: Dimension + codomain: Dimension + + @property + def tag(cls) -> Tag: + """The connectivity's identity: its qualified Python name.""" + return f"{cls.__module__}.{cls.__qualname__}" + + def __call__(cls, *args: Any, **kwargs: Any) -> NoReturn: + raise TypeError( + f"'{cls.__qualname__}' is a connectivity declaration and cannot be instantiated;" + " bind a neighbor table to it through the offset provider." + ) + + @overload + def __getitem__(cls, item: int) -> Connectivity: ... + @overload + def __getitem__(cls, item: Any) -> Any: ... + def __getitem__(cls, item: Any) -> Any: + # NOTE: `numbers.Integral`, not `int`, so `V2E[np.int32(1)]` does not fall through to + # type-parameter subscription; `bool` is excluded so `V2E[True]` is an error. + if isinstance(item, numbers.Integral) and not isinstance(item, bool): + return cls.__gt_field_offset__()[int(item)] + # A metaclass `__getitem__` shadows `__class_getitem__`, so type-parameter + # subscription (`NeighborConnectivity[V, E]`) has to be forwarded explicitly. + return cast(Any, cls).__class_getitem__(item) + + def __repr__(cls) -> str: + return cls.tag + + def __str__(cls) -> str: + return cls.__qualname__ + + def __gt_type__(cls) -> Any: + return cls.__gt_field_offset__().__gt_type__() + + def __gt_field_offset__(cls) -> Any: + """ + The `FieldOffset` equivalent to this connectivity. + + Its tag is the *local dimension's* tag, the single string that shifts, neighbor + reductions and sparse arguments all use to find the table in the offset provider. + """ + from gt4py.next.ffront import fbuiltins + + if (field_offset := cls.__dict__.get("_field_offset")) is None: + field_offset = fbuiltins.FieldOffset( + cls.Local.tag, source=cls.codomain, target=(cls.origin, cls.Local) + ) + type.__setattr__(cls, "_field_offset", field_offset) + return field_offset + + +class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + """ + Declare a neighbor connectivity: for each `Origin` element, a list of `Codomain` neighbors. + + The declaration names the connectivity's local dimension -- its nested `Local` class -- + and optionally its neighbor counts. It holds no data: the neighbor table is bound at call + time through the offset provider, and checked against the declaration. + + Examples: + >>> class Vertex(DimensionIndex): ... + >>> class Edge(DimensionIndex): ... + >>> class V2E(NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): + ... class Local(LocalDimensionIndex): ... + >>> V2E.origin is Vertex, V2E.codomain is Edge + (True, True) + >>> V2E.Local.owner is V2E, V2E.Local.max_neighbors, V2E.Local.min_neighbors + (True, 6, 5) + """ + + # NOTE: an annotation only, never assigned here. A real nested class on the base would be + # flagged by pyright as an incompatible override in every declaration. + Local: ClassVar[type[LocalDimensionIndex]] + origin: ClassVar[Dimension] + codomain: ClassVar[Dimension] + + def __init_subclass__( + cls, + /, + *, + max_neighbors: Optional[int] = None, + min_neighbors: Optional[int] = None, + **kwargs: Any, + ) -> None: + super().__init_subclass__(**kwargs) + name = cls.__qualname__ + if "" in name: + raise TypeError( + f"'{name}' must be declared at module level: a connectivity is referenced from" + " the IR by its qualified name, which has to be importable." + ) + params = [ + xtyping.get_args(base) + for base in cls.__dict__.get("__orig_bases__", ()) + if xtyping.get_origin(base) is NeighborConnectivity + ] + if len(params) != 1 or len(params[0]) != 2: + raise TypeError( + f"'{name}' must derive from 'NeighborConnectivity[Origin, Codomain]' directly," + " with both dimensions given." + ) + origin, codomain = params[0] + for role, dim in (("Origin", origin), ("Codomain", codomain)): + if not isinstance(dim, DimensionMeta) or dim.kind is DimensionKind.LOCAL: + raise TypeError(f"'{name}': '{role}' must be a non-local dimension, got '{dim}'.") + + local = cls.__dict__.get("Local") + if not (isinstance(local, DimensionMeta) and issubclass(local, LocalDimensionIndex)): + raise TypeError( + f"'{name}' must declare its local dimension as a nested class:" + " 'class Local(LocalDimensionIndex): ...'." + ) + if local.owner is not None: + raise TypeError( + f"'{name}': '{local.__qualname__}' is already the local dimension of" + f" '{local.owner.__qualname__}'." + ) + + max_neighbors = _check_neighbor_count(cls, "max_neighbors", max_neighbors) + min_neighbors = _check_neighbor_count(cls, "min_neighbors", min_neighbors) + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("min_neighbors", min_neighbors), + ): + declared = getattr(local, count_name) + if count is not None and declared is not None and count != declared: + raise TypeError( + f"'{name}': '{count_name}={count}' contradicts the size declared by" + f" '{local.__qualname__}' ({declared})." + ) + max_neighbors = max_neighbors if max_neighbors is not None else local.max_neighbors + min_neighbors = min_neighbors if min_neighbors is not None else local.min_neighbors + if ( + max_neighbors is not None + and min_neighbors is not None + and min_neighbors > max_neighbors + ): + raise TypeError( + f"'{name}': 'min_neighbors' ({min_neighbors}) exceeds 'max_neighbors'" + f" ({max_neighbors})." + ) + + cls.origin, cls.codomain = origin, codomain + local.owner = cls + local.max_neighbors, local.min_neighbors = max_neighbors, min_neighbors + + +def check_neighbor_table( + connectivity: type[NeighborConnectivity], + table: NeighborTable | NeighborConnectivityType, +) -> None: + """ + Check that a neighbor table matches the connectivity declaration it is bound to. + + Args: + connectivity: The declaration. + table: The bound table, or its type (which is all an ahead-of-time compilation has). + + Raises: + ValueError: On the first mismatch, naming the connectivity and the mismatch. + """ + table_type = table if isinstance(table, NeighborConnectivityType) else table.__gt_type__() + name = connectivity.__qualname__ + local = connectivity.Local + + def fail(reason: str) -> NoReturn: + raise ValueError(f"The table bound to '{name}' does not match its declaration: {reason}.") + + if not isinstance(table_type, NeighborConnectivityType): + fail(f"expected a neighbor table, got '{table_type}'") + expected_domain = (connectivity.origin, local) + if tuple(table_type.domain) != expected_domain: + fail( + f"its domain is '({', '.join(map(str, table_type.domain))})'," + f" expected '({', '.join(map(str, expected_domain))})'" + ) + if table_type.codomain is not connectivity.codomain: + fail(f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'") + if table_type.dtype.kind not in (core_defs.DTypeKind.INT, core_defs.DTypeKind.UINT): + fail(f"its dtype '{table_type.dtype}' is not integral") + if local.max_neighbors is not None and table_type.max_neighbors != local.max_neighbors: + fail( + f"it has {table_type.max_neighbors} neighbors per element," + f" expected max_neighbors={local.max_neighbors}" + ) + if local.min_neighbors is not None: + max_neighbors = table_type.max_neighbors + if local.min_neighbors < max_neighbors and not table_type.has_skip_values: + fail( + f"min_neighbors={local.min_neighbors} < {max_neighbors} requires a skip value," + " but the table has none" + ) + if local.min_neighbors == max_neighbors and table_type.has_skip_values: + fail( + f"min_neighbors == max_neighbors == {max_neighbors} means every element has all" + f" its neighbors, but the table has skip value {table_type.skip_value}" + ) diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index b8133cbafa..38e760995c 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -230,7 +230,9 @@ def as_scalar(self) -> core_defs.ScalarT: def premap( self: NdArrayField, - *connectivities: common.Connectivity | fbuiltins.FieldOffset, + *connectivities: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], ) -> NdArrayField: """ Rearrange the field content using the provided connectivities (index mappings). @@ -305,7 +307,10 @@ def premap( codomains_counter: collections.Counter[common.Dimension] = collections.Counter() for connectivity in connectivities: - # For neighbor reductions, a FieldOffset is passed instead of an actual Connectivity + # For neighbor reductions, a FieldOffset or a connectivity declaration is passed + # instead of an actual Connectivity + if isinstance(connectivity, common.ConnectivityMeta): + connectivity = connectivity.__gt_field_offset__() if not isinstance(connectivity, common.Connectivity): assert isinstance(connectivity, fbuiltins.FieldOffset) connectivity = connectivity.as_connectivity_field() @@ -357,8 +362,10 @@ def premap( def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: return functools.reduce( lambda field, current_index_field: field.premap(current_index_field), diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 9e8f1fa0f1..461a52f6ed 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -434,11 +434,15 @@ def visit_Symbol( def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> foast.Attribute: new_value = self.visit(node.value, **kwargs) + match new_value.type: + # `V2E.Local`: the local dimension of a connectivity declaration, which is the last + # target of the offset it is typed as. + case ts.OffsetType(target=(_, local)) if node.attr == "Local": + attr_type: ts.TypeSpec = ts.DimensionType(dim=local) + case _: + attr_type = getattr(new_value.type, node.attr) return foast.Attribute( - value=new_value, - attr=node.attr, - location=node.location, - type=getattr(new_value.type, node.attr), + value=new_value, attr=node.attr, location=node.location, type=attr_type ) def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscript: diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 32b9f9dfac..d82ed6e201 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -74,7 +74,7 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> stages.CompilableProgramDef: """ all_closure_vars = transform_utils._get_closure_vars_recursively(inp.data.closure_vars) offsets_and_dimensions = transform_utils._filter_closure_vars_by_type( - all_closure_vars, fbuiltins.FieldOffset, common.DimensionMeta + all_closure_vars, fbuiltins.FieldOffset, common.ConnectivityMeta, common.DimensionMeta ) grid_type = transform_utils._deduce_grid_type( inp.data.grid_type, offsets_and_dimensions.values() diff --git a/src/gt4py/next/ffront/transform_utils.py b/src/gt4py/next/ffront/transform_utils.py index 09c9d4b9ee..4e24d83881 100644 --- a/src/gt4py/next/ffront/transform_utils.py +++ b/src/gt4py/next/ffront/transform_utils.py @@ -47,7 +47,9 @@ def _filter_closure_vars_by_type(closure_vars: dict[str, Any], *types: type) -> def _deduce_grid_type( requested_grid_type: Optional[common.GridType], - offsets_and_dimensions: Iterable[fbuiltins.FieldOffset | common.Dimension], + offsets_and_dimensions: Iterable[ + fbuiltins.FieldOffset | type[common.NeighborConnectivity] | common.Dimension + ], ) -> common.GridType: """ Derive grid type from actually occurring dimensions and check against optional user request. @@ -59,7 +61,9 @@ def _deduce_grid_type( deduced_grid_type = common.GridType.CARTESIAN for o in offsets_and_dimensions: - if isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o): + if isinstance(o, common.ConnectivityMeta) or ( + isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o) + ): deduced_grid_type = common.GridType.UNSTRUCTURED break if isinstance(o, common.DimensionMeta) and o.kind == common.DimensionKind.LOCAL: diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 7b37e35684..b0bbf16f63 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -151,7 +151,12 @@ def ndarray(self) -> core_defs.NDArrayObject: def asnumpy(self) -> np.ndarray: raise NotImplementedError - def premap(self, index_field: common.Connectivity | fbuiltins.FieldOffset) -> common.Field: + def premap( + self, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + ) -> common.Field: raise NotImplementedError def restrict( # type: ignore[override] @@ -168,8 +173,10 @@ def as_scalar(self) -> typing.Never: def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError() @@ -1143,8 +1150,10 @@ def as_scalar(self) -> core_defs.IntegralScalar: def premap( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1284,8 +1293,10 @@ def asnumpy(self) -> np.ndarray: def premap( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py new file mode 100644 index 0000000000..f8167f2627 --- /dev/null +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py @@ -0,0 +1,104 @@ +# 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 + +"""A `NeighborConnectivity` declaration used directly in DSL code, on every backend.""" + +import numpy as np +import pytest + +import gt4py.next as gtx +from gt4py.next import Dims, Field, common, constructors, neighbor_sum + +from next_tests import definitions as test_defs +from next_tests.integration_tests import cases, cases_utils +from next_tests.integration_tests.cases_utils import ( # noqa: F401 [unused-import] # fixture + exec_alloc_descriptor, +) + + +class V(gtx.DimensionIndex): ... + + +class E(gtx.DimensionIndex): ... + + +class V2E(gtx.NeighborConnectivity[V, E], max_neighbors=4, min_neighbors=4): + class Local(gtx.LocalDimensionIndex): ... + + +@pytest.fixture +def case(exec_alloc_descriptor): + mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) + v2e_arr = mesh.offset_provider[cases_utils.V2EDim.tag].asnumpy() + table = constructors.as_connectivity( + domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, + codomain=E, + data=v2e_arr, + skip_value=None, + allocator=exec_alloc_descriptor.allocator, + ) + common.check_neighbor_table(V2E, table) + return cases.Case( + ( + None + if isinstance(exec_alloc_descriptor, test_defs.EmbeddedDummyBackend) + else exec_alloc_descriptor + ), + # NOTE: still keyed on the local dimension's tag; class keys come with the removal of + # `FieldOffset`. + offset_provider={V2E.Local.tag: table}, + default_sizes={V: mesh.num_vertices, E: mesh.num_edges, V2E.Local: v2e_arr.shape[1]}, + grid_type=common.GridType.UNSTRUCTURED, + allocator=exec_alloc_descriptor.allocator, + ) + + +def _table(case: cases.Case) -> np.ndarray: + return case.offset_provider[V2E.Local.tag].asnumpy() + + +@pytest.mark.uses_unstructured_shift +def test_shift(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2E[1]) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case)[:, 1]]) + + +@pytest.mark.uses_unstructured_shift +def test_reduction(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return neighbor_sum(a(V2E), axis=V2E.Local) + + cases.verify_with_default_data(case, testee, lambda a: np.sum(a[_table(case)], axis=1)) + + +@pytest.mark.uses_unstructured_shift +def test_sparse_argument(case): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + return neighbor_sum(s * a(V2E), axis=V2E.Local) + + cases.verify_with_default_data(case, testee, lambda s, a: np.sum(s * a[_table(case)], axis=1)) + + +@pytest.mark.uses_unstructured_shift +def test_program(case): + @gtx.field_operator + def shift_by_one(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2E[0]) + + @gtx.program + def testee(a: Field[Dims[E], float], out: Field[Dims[V], float]): + shift_by_one(a, out=out) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case)[:, 0]]) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py new file mode 100644 index 0000000000..f661140ce3 --- /dev/null +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -0,0 +1,314 @@ +# 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 + +import pickle +import textwrap +import typing + +import numpy as np +import pytest + +from gt4py._core import definitions as core_defs +from gt4py.next import common +from gt4py.next.common import ( + DimensionIndex, + DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, + NeighborConnectivityType, +) +from gt4py.next.ffront import transform_utils +from gt4py.next.type_system import type_specifications as ts, type_translation + + +class Vertex(DimensionIndex): ... + + +class Edge(DimensionIndex): ... + + +class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... + + +class V2E(NeighborConnectivity[Vertex, Edge], max_neighbors=4, min_neighbors=3): + class Local(LocalDimensionIndex): ... + + +class E2V(NeighborConnectivity[Edge, Vertex]): + class Local(LocalDimensionIndex): ... + + +class LsqCoeff(LocalDimensionIndex, size=3): ... + + +def _declare(source: str) -> dict: + """ + Run `source` as the body of a throwaway module. + + Declarations have to be at module level, so error cases cannot simply be written inside + the test function: the `` check would fire before the one under test. + """ + namespace = { + "__name__": __name__, + "DimensionIndex": DimensionIndex, + "DimensionKind": DimensionKind, + "LocalDimensionIndex": LocalDimensionIndex, + "NeighborConnectivity": NeighborConnectivity, + "Vertex": Vertex, + "Edge": Edge, + "KDim": KDim, + "V2E": V2E, + } + exec(textwrap.dedent(source), namespace) + return namespace + + +class TestDeclaration: + def test_owner_and_dimensions(self): + assert V2E.Local.owner is V2E + assert V2E.origin is Vertex + assert V2E.codomain is Edge + assert V2E.Local.kind is DimensionKind.LOCAL + assert issubclass(V2E.Local, DimensionIndex) + + def test_counts(self): + assert (V2E.Local.max_neighbors, V2E.Local.min_neighbors) == (4, 3) + assert (E2V.Local.max_neighbors, E2V.Local.min_neighbors) == (None, None) + + def test_ownerless_local(self): + assert LsqCoeff.owner is None + assert (LsqCoeff.max_neighbors, LsqCoeff.min_neighbors) == (3, 3) + assert LsqCoeff.kind is DimensionKind.LOCAL + + def test_counts_from_local_size(self): + ns = _declare( + """ + class C2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex, size=3): ... + """ + ) + assert (ns["C2E"].Local.max_neighbors, ns["C2E"].Local.min_neighbors) == (3, 3) + + def test_identity(self): + assert V2E.tag == f"{__name__}.V2E" + assert V2E.Local.tag == f"{__name__}.V2E.Local" + assert common.resolve(V2E.Local.tag) is V2E.Local + assert str(V2E) == "V2E" + assert repr(V2E) == V2E.tag + + def test_pickle_by_reference(self): + assert pickle.loads(pickle.dumps(V2E)) is V2E + assert pickle.loads(pickle.dumps(V2E.Local)) is V2E.Local + + def test_hashable(self): + assert {V2E: 1}[V2E] == 1 + + def test_type_parameter_subscription(self): + alias = NeighborConnectivity[Vertex, Edge] + assert typing.get_origin(alias) is NeighborConnectivity + assert typing.get_args(alias) == (Vertex, Edge) + + def test_bool_is_not_a_neighbor_index(self): + with pytest.raises(TypeError): + V2E[True] + + def test_not_instantiable(self): + with pytest.raises(TypeError, match="cannot be instantiated"): + V2E() + + +class TestDeclarationErrors: + @pytest.mark.parametrize( + "source, match", + [ + ( + """ + class C(NeighborConnectivity[Vertex, Edge]): ... + """, + "must declare its local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge]): + class Local(DimensionIndex, kind=DimensionKind.LOCAL): ... + """, + "must declare its local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge]): + Local = V2E.Local + """, + "already the local dimension of 'V2E'", + ), + ( + """ + class C(NeighborConnectivity): + class Local(LocalDimensionIndex): ... + """, + "must derive from 'NeighborConnectivity\\[Origin, Codomain\\]'", + ), + ( + """ + class C(V2E): + class Local(LocalDimensionIndex): ... + """, + "must derive from 'NeighborConnectivity\\[Origin, Codomain\\]'", + ), + ( + """ + class C(NeighborConnectivity[V2E.Local, Edge]): + class Local(LocalDimensionIndex): ... + """, + "'Origin' must be a non-local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, int]): + class Local(LocalDimensionIndex): ... + """, + "'Codomain' must be a non-local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=2, min_neighbors=3): + class Local(LocalDimensionIndex): ... + """, + "exceeds 'max_neighbors'", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=4): + class Local(LocalDimensionIndex, size=3): ... + """, + "contradicts the size", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=-1): + class Local(LocalDimensionIndex): ... + """, + "non-negative integer", + ), + ( + """ + class L(LocalDimensionIndex, kind=DimensionKind.HORIZONTAL): ... + """, + "cannot have kind", + ), + ( + """ + class L(LocalDimensionIndex, size=1.5): ... + """, + "non-negative integer", + ), + ], + ) + def test_rejected(self, source, match): + with pytest.raises(TypeError, match=match): + _declare(source) + + def test_function_local_declaration(self): + with pytest.raises(TypeError, match="module level"): + + class C(NeighborConnectivity[Vertex, Edge]): + Local = LsqCoeff + + def test_subclass_of_owned_local_is_ownerless(self): + ns = _declare( + """ + class Other(V2E.Local): ... + """ + ) + assert ns["Other"].owner is None + + def test_local_dimension_cannot_be_staggered(self): + with pytest.raises(TypeError, match="cannot be staggered"): + common.Staggered[V2E.Local] + + +def _table_type( + domain=(Vertex, V2E.Local), + codomain=Edge, + max_neighbors=4, + skip_value=common._DEFAULT_SKIP_VALUE, + dtype=np.int32, +) -> NeighborConnectivityType: + return NeighborConnectivityType( + domain=domain, + codomain=codomain, + skip_value=skip_value, + dtype=core_defs.dtype(dtype), + max_neighbors=max_neighbors, + ) + + +class TestCheckNeighborTable: + def test_matching_type(self): + common.check_neighbor_table(V2E, _table_type()) + + def test_matching_table(self): + from gt4py.next import constructors + + table = constructors.as_connectivity( + domain={Edge: 2, E2V.Local: 2}, codomain=Vertex, data=np.array([[0, 1], [1, 2]]) + ) + common.check_neighbor_table(E2V, table) + + def test_undeclared_counts_accept_any_table(self): + common.check_neighbor_table( + E2V, _table_type(domain=(Edge, E2V.Local), codomain=Vertex, max_neighbors=7) + ) + + @pytest.mark.parametrize( + "kwargs, match", + [ + ({"domain": (Vertex, E2V.Local)}, "its domain is"), + ({"domain": (Edge, V2E.Local)}, "its domain is"), + ({"codomain": Vertex}, "its codomain is"), + ({"dtype": np.float64}, "is not integral"), + ({"max_neighbors": 5}, "expected max_neighbors=4"), + ({"skip_value": None}, "requires a skip value"), + ], + ) + def test_mismatch(self, kwargs, match): + with pytest.raises(ValueError, match=match): + common.check_neighbor_table(V2E, _table_type(**kwargs)) + + def test_skip_value_without_missing_neighbors(self): + ns = _declare( + """ + class Full(NeighborConnectivity[Vertex, Edge], max_neighbors=2, min_neighbors=2): + class Local(LocalDimensionIndex): ... + """ + ) + full = ns["Full"] + with pytest.raises(ValueError, match="has skip value"): + common.check_neighbor_table( + full, _table_type(domain=(Vertex, full.Local), max_neighbors=2) + ) + + +class TestFrontendIntegration: + def test_from_value_is_an_offset(self): + # NOTE: pins the `__gt_type__` branch of `from_value` ahead of the dimension branch; a + # connectivity declaration is a class, like a dimension. + assert type_translation.from_value(V2E) == ts.OffsetType( + source=Edge, target=(Vertex, V2E.Local), tag=V2E.Local.tag + ) + + def test_field_offset_is_derived_once(self): + assert V2E.__gt_field_offset__() is V2E.__gt_field_offset__() + assert V2E.__gt_field_offset__().value == V2E.Local.tag + + def test_grid_type_deduction(self): + assert ( + transform_utils._deduce_grid_type(None, [Vertex, V2E]) is common.GridType.UNSTRUCTURED + ) + with pytest.raises(ValueError, match="CARTESIAN"): + transform_utils._deduce_grid_type(common.GridType.CARTESIAN, [V2E]) diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 1506c46a75..04007c4a32 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -302,3 +302,72 @@ main:16:48: error: Type argument "C" of "Staggered" must be a subtype of "CartesianAxisIndex" [type-var] main:17:13: error: Unsupported operand types for + ("type[C]" and "int") [operator] main:18:18: error: Unsupported operand types for - ("type[C]" and "int") [operator] + + - case: neighbor_connectivity_declaration + main: | + from __future__ import annotations + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6): + class Local(gtx.LocalDimensionIndex): ... + + def sparse(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + reveal_type(V2E.Local) + reveal_type(V2E.Local.owner) + reveal_type(V2E[1]) + out: | + main:12:13: note: Revealed type is "def (value: int) -> main.V2E.Local" + main:13:13: note: Revealed type is "type[gt4py.next.common.NeighborConnectivity[Any, Any]] | None" + main:14:13: note: Revealed type is "gt4py.next.common.Connectivity[Any, Any]" + + - case: neighbor_connectivity_locals_are_distinct + main: | + from __future__ import annotations + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + class Cell(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + class V2C(gtx.NeighborConnectivity[Vertex, Cell]): + class Local(gtx.LocalDimensionIndex): ... + + def takes_v2e(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + def caller(a: gtx.Field[gtx.Dims[Vertex, V2C.Local], gtx.float64]) -> None: + takes_v2e(a) + out: | + main:17:15: error: Argument 1 to "takes_v2e" has incompatible type "Field[Dims[Vertex, main.V2C.Local], float]"; expected "Field[Dims[Vertex, main.V2E.Local], float]" [arg-type] + main:17:15: note: "Field[Dims[Vertex, Local], float].__call__" has type "def __call__(self, index_field: Connectivity[Any, Any] | FieldOffset | type[NeighborConnectivity[Any, Any]], *args: Connectivity[Any, Any] | FieldOffset | type[NeighborConnectivity[Any, Any]]) -> Field[Any, Any]" + + - case: neighbor_connectivity_generic_local + main: | + from __future__ import annotations + import typing + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + L = typing.TypeVar("L", bound=gtx.LocalDimensionIndex) + + def local_of(conn: type[gtx.NeighborConnectivity]) -> type[gtx.LocalDimensionIndex]: + return conn.Local + + def first(a: gtx.Field[gtx.Dims[Vertex, L], gtx.float64]) -> type[L]: + raise NotImplementedError + + def caller(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: + reveal_type(first(a)) + out: | + main:20:17: note: Revealed type is "type[main.V2E.Local]" From c1f39dcd207ed43a79ac2038512e48583aa7ba2f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 00:28:02 +0200 Subject: [PATCH 02/13] fix[next]: address review of NeighborConnectivity declarations - check_neighbor_table: min_neighbors exceeding the table, bool tables - descriptive errors for non-integer indices, the undeclared base, DSL attributes - FieldOffset.Local, so the spelling works for legacy offsets in embedded too - fingerprint a declaration by its dimensions and counts, not only its name - negative counts are a ValueError --- .../ADRs/next/0029-Connectivities_As_Types.md | 23 +++- src/gt4py/next/common.py | 33 +++-- src/gt4py/next/ffront/fbuiltins.py | 9 ++ .../ffront/foast_passes/type_deduction.py | 7 +- src/gt4py/next/fingerprinting.py | 17 +++ .../unit_tests/test_neighbor_connectivity.py | 118 ++++++++++++++++-- 6 files changed, 185 insertions(+), 22 deletions(-) diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index 2fdcc101e8..45898fd8d7 100644 --- a/docs/development/ADRs/next/0029-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -23,8 +23,7 @@ def f(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: ``` The declaration is written in DSL code, owns its local dimension, and states the -constraints the neighbor table bound at call time has to satisfy. It holds no -data. It builds on [ADR 0028](0028-Dimensions_As_Nominal_Types.md): the +constraints a neighbor table bound to it has to satisfy. It holds no data. It builds on [ADR 0028](0028-Dimensions_As_Nominal_Types.md): the connectivity, like a dimension, is identified by its type, and `V2E.Local` is an ordinary dimension class with the tag `.V2E.Local`. @@ -62,7 +61,12 @@ and skip values were never checked against the `FieldOffset` declaration. - `common.check_neighbor_table(V2E, table)` checks a table, or just its type (which is all an ahead-of-time compilation has), against the declaration: the domain is `(Origin, V2E.Local)`, the codomain is `Codomain`, the dtype is - integral, and the neighbor counts and skip values agree. + integral, and the neighbor counts and skip values agree. Skip values are + checked on the table's type: a table with a `skip_value` counts as having skip + values whether or not an entry uses it. The check is explicit for as long as + offset providers are keyed by tag strings, since nothing then connects a + provider entry to a declaration; it becomes automatic with class-keyed + providers. ### `NeighborConnectivity` is not a `Connectivity` @@ -98,7 +102,11 @@ is the `ts.OffsetType` of the derived offset `(Codomain -> (Origin, Local))`, whose tag is **the local dimension's tag**, `V2E.Local.tag`. This is the single string that shifts, neighbor reductions and sparse arguments already use to find the table in the offset provider, so existing backends need no change. -`V2E.Local` inside DSL code types as that local dimension. `V2E[i]` subscripts the +`V2E.Local` inside DSL code types as that local dimension, and +`FieldOffset.Local` names the same thing on a legacy offset, so the spelling +works for both. The other frontend touch points treat the class like the +`FieldOffset` it derives: grid-type deduction (`transform_utils`, `past_to_itir`) +counts it as unstructured, and embedded `premap` accepts it. `V2E[i]` subscripts the metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. ## Consequences @@ -106,8 +114,11 @@ metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, - An unstructured connectivity is spelled once. The provider key, the offset tag and the local dimension are all derived from the declaration. - A table bound to a connectivity can be checked against its declaration. -- The frontend needs one special case: `V2E.Local` is resolved from the offset - type, because the type of `V2E` is not the class. +- `V2E.Local` in DSL code is resolved from the offset type, because the type of + `V2E` is not the class. +- A declaration is fingerprinted by its name *and* its declared dimensions and + counts, so redefining it under the same name (e.g. re-running a notebook + cell) does not reuse artifacts compiled for the old declaration. - `FieldOffset` remains during migration; a `FieldOffset` and a `NeighborConnectivity` sharing a local dimension are interchangeable. diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index dc4b9eda7c..341b11160a 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -2014,13 +2014,13 @@ def __init_subclass__( cls.max_neighbors = cls.min_neighbors = _check_neighbor_count(cls, "size", size) -def _check_neighbor_count(owner: type, name: str, count: Optional[int]) -> Optional[int]: +def _check_neighbor_count(cls: type, name: str, count: Optional[int]) -> Optional[int]: if count is None: return None - if not isinstance(count, numbers.Integral) or isinstance(count, bool) or count < 0: - raise TypeError( - f"'{owner.__qualname__}': '{name}' must be a non-negative integer, got '{count!r}'." - ) + if not isinstance(count, numbers.Integral) or isinstance(count, bool): + raise TypeError(f"'{cls.__qualname__}': '{name}' must be an integer, got '{count!r}'.") + if count < 0: + raise ValueError(f"'{cls.__qualname__}': '{name}' must be non-negative, got {count}.") return int(count) @@ -2056,6 +2056,11 @@ def __getitem__(cls, item: Any) -> Any: # type-parameter subscription; `bool` is excluded so `V2E[True]` is an error. if isinstance(item, numbers.Integral) and not isinstance(item, bool): return cls.__gt_field_offset__()[int(item)] + if "Local" in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}[{item!r}]': a connectivity is indexed by an integer" + " neighbor position." + ) # A metaclass `__getitem__` shadows `__class_getitem__`, so type-parameter # subscription (`NeighborConnectivity[V, E]`) has to be forwarded explicitly. return cast(Any, cls).__class_getitem__(item) @@ -2078,6 +2083,11 @@ def __gt_field_offset__(cls) -> Any: """ from gt4py.next.ffront import fbuiltins + if "Local" not in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}' is not a connectivity declaration; declare one by" + " subclassing 'NeighborConnectivity[Origin, Codomain]'." + ) if (field_offset := cls.__dict__.get("_field_offset")) is None: field_offset = fbuiltins.FieldOffset( cls.Local.tag, source=cls.codomain, target=(cls.origin, cls.Local) @@ -2094,7 +2104,8 @@ class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( The declaration names the connectivity's local dimension -- its nested `Local` class -- and optionally its neighbor counts. It holds no data: the neighbor table is bound at call - time through the offset provider, and checked against the declaration. + time through the offset provider. `check_neighbor_table` checks a table against the + declaration. Examples: >>> class Vertex(DimensionIndex): ... @@ -2191,6 +2202,9 @@ def check_neighbor_table( """ Check that a neighbor table matches the connectivity declaration it is bound to. + Skip values are checked on the table's *type*: a table with a `skip_value` counts as + having skip values whether or not any entry uses it. + Args: connectivity: The declaration. table: The bound table, or its type (which is all an ahead-of-time compilation has). @@ -2215,7 +2229,7 @@ def fail(reason: str) -> NoReturn: ) if table_type.codomain is not connectivity.codomain: fail(f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'") - if table_type.dtype.kind not in (core_defs.DTypeKind.INT, core_defs.DTypeKind.UINT): + if not np.issubdtype(table_type.dtype.scalar_type, np.integer): fail(f"its dtype '{table_type.dtype}' is not integral") if local.max_neighbors is not None and table_type.max_neighbors != local.max_neighbors: fail( @@ -2224,6 +2238,11 @@ def fail(reason: str) -> NoReturn: ) if local.min_neighbors is not None: max_neighbors = table_type.max_neighbors + if local.min_neighbors > max_neighbors: + fail( + f"min_neighbors={local.min_neighbors} exceeds its {max_neighbors} neighbors" + " per element" + ) if local.min_neighbors < max_neighbors and not table_type.has_skip_values: fail( f"min_neighbors={local.min_neighbors} < {max_neighbors} requires a skip value," diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 7474cd1406..942a5dcc10 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -496,6 +496,15 @@ def __post_init__(self) -> None: def __gt_type__(self) -> ts.OffsetType: return ts.OffsetType(source=self.source, target=self.target, tag=self.value) + @property + def Local(self) -> common.Dimension: + """The local dimension, as `V2E.Local` names it on a `NeighborConnectivity`.""" + if len(self.target) != 2: + raise AttributeError( + f"'{self.value}' is a Cartesian offset and has no local dimension." + ) + return self.target[1] + def __getitem__(self, offset: int) -> common.Connectivity: """Serve as a connectivity factory.""" from gt4py.next import embedded # avoid circular import diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 461a52f6ed..ba62641dc5 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -440,7 +440,12 @@ def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> foast.Attribu case ts.OffsetType(target=(_, local)) if node.attr == "Local": attr_type: ts.TypeSpec = ts.DimensionType(dim=local) case _: - attr_type = getattr(new_value.type, node.attr) + try: + attr_type = getattr(new_value.type, node.attr) + except AttributeError: + raise errors.DSLError( + node.location, f"'{new_value.type}' has no attribute '{node.attr}'." + ) from None return foast.Attribute( value=new_value, attr=node.attr, location=node.location, type=attr_type ) diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index 52a7081b1d..a9d8e206b3 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -231,6 +231,23 @@ def _instance_state(obj: Any) -> tuple[Any, ...]: if "base" in obj.__dict__ else EmptyDeconstruction.from_reference(obj) ), + # A connectivity declaration is a class, fingerprinted by reference like a dimension, *and* + # by what it declares: redefining `V2E` under the same name with other dimensions or counts + # (e.g. re-running a notebook cell) must not reuse artifacts built for the old declaration. + common.ConnectivityMeta: lambda obj: ( + # NOTE: the name goes into the state unverified; importability is still enforced by the + # strict fingerprinter through `obj.Local`, which is a class nested in `obj`. + Deconstruction.from_pieces( + obj.origin, + obj.codomain, + obj.Local, + obj.Local.max_neighbors, + obj.Local.min_neighbors, + state=b"neighbor_connectivity\0" + obj.tag.encode(), + ) + if "Local" in obj.__dict__ + else EmptyDeconstruction.from_reference(obj) + ), type(None): lambda obj: EmptyDeconstruction.from_typed_value(type(None)), bool: lambda obj: EmptyDeconstruction.from_typed_value(bool, b"1" if obj else b"0"), int: lambda obj: EmptyDeconstruction.from_typed_value(type(obj), str(int(obj)).encode()), diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index f661140ce3..f0705332cd 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -188,13 +188,6 @@ class Local(LocalDimensionIndex, size=3): ... """, "contradicts the size", ), - ( - """ - class C(NeighborConnectivity[Vertex, Edge], max_neighbors=-1): - class Local(LocalDimensionIndex): ... - """, - "non-negative integer", - ), ( """ class L(LocalDimensionIndex, kind=DimensionKind.HORIZONTAL): ... @@ -205,7 +198,7 @@ class L(LocalDimensionIndex, kind=DimensionKind.HORIZONTAL): ... """ class L(LocalDimensionIndex, size=1.5): ... """, - "non-negative integer", + "must be an integer", ), ], ) @@ -213,6 +206,35 @@ def test_rejected(self, source, match): with pytest.raises(TypeError, match=match): _declare(source) + def test_negative_count(self): + with pytest.raises(ValueError, match="non-negative"): + _declare( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=-1): + class Local(LocalDimensionIndex): ... + """ + ) + + def test_adopting_an_ownerless_local(self): + ns = _declare( + """ + class Coeff(LocalDimensionIndex, size=3): ... + + class C(NeighborConnectivity[Vertex, Edge]): + Local = Coeff + """ + ) + assert ns["Coeff"].owner is ns["C"] + assert (ns["Coeff"].max_neighbors, ns["Coeff"].min_neighbors) == (3, 3) + + def test_non_integer_index(self): + with pytest.raises(TypeError, match="indexed by an integer"): + V2E[Vertex] + + def test_base_is_not_a_declaration(self): + with pytest.raises(TypeError, match="not a connectivity declaration"): + NeighborConnectivity.__gt_type__() + def test_function_local_declaration(self): with pytest.raises(TypeError, match="module level"): @@ -280,6 +302,31 @@ def test_mismatch(self, kwargs, match): with pytest.raises(ValueError, match=match): common.check_neighbor_table(V2E, _table_type(**kwargs)) + def test_min_neighbors_exceeds_table(self): + ns = _declare( + """ + class MinOnly(NeighborConnectivity[Vertex, Edge], min_neighbors=5): + class Local(LocalDimensionIndex): ... + """ + ) + min_only = ns["MinOnly"] + for skip_value in (None, common._DEFAULT_SKIP_VALUE): + with pytest.raises(ValueError, match="min_neighbors=5 exceeds"): + common.check_neighbor_table( + min_only, + _table_type( + domain=(Vertex, min_only.Local), max_neighbors=3, skip_value=skip_value + ), + ) + + def test_bool_table_is_not_integral(self): + with pytest.raises(ValueError, match="is not integral"): + common.check_neighbor_table(V2E, _table_type(dtype=bool)) + + def test_not_a_neighbor_table(self): + with pytest.raises(ValueError, match="expected a neighbor table"): + common.check_neighbor_table(V2E, common.CartesianConnectivity(Vertex, 1)) + def test_skip_value_without_missing_neighbors(self): ns = _declare( """ @@ -306,6 +353,61 @@ def test_field_offset_is_derived_once(self): assert V2E.__gt_field_offset__() is V2E.__gt_field_offset__() assert V2E.__gt_field_offset__().value == V2E.Local.tag + def test_neighbor_index_accepts_numpy_integers(self): + from gt4py.next import constructors, embedded + + table = constructors.as_connectivity( + domain={Vertex: 2, V2E.Local: 4}, + codomain=Edge, + data=np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), + ) + with embedded.context.update(offset_provider={V2E.Local.tag: table}): + assert np.array_equal(V2E[np.int32(1)].asnumpy(), V2E[1].asnumpy()) + + def test_legacy_field_offset_has_local(self): + from gt4py.next import FieldOffset + + assert FieldOffset("V2E", source=Edge, target=(Vertex, V2E.Local)).Local is V2E.Local + + def test_attribute_errors_are_dsl_errors(self): + from gt4py.next import errors, field_operator + from gt4py.next.ffront.func_to_foast import FieldOperatorParser + from gt4py.next import Dims, Field + + def origin_of(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return a(V2E.origin) + + with pytest.raises(errors.DSLError, match="has no attribute 'origin'"): + FieldOperatorParser.apply_to_function(origin_of) + + def test_fingerprint_covers_the_declaration(self): + from gt4py.next import fingerprinting + + def fingerprint_of(source: str) -> str: + # lenient: `_declare` classes are not importable, as in a re-run notebook cell + return fingerprinting.lenient_fingerprinter(_declare(source)["C"]) + + base = fingerprint_of( + """ + class C(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + """ + ) + swapped = fingerprint_of( + """ + class C(NeighborConnectivity[Edge, Vertex]): + class Local(LocalDimensionIndex): ... + """ + ) + counted = fingerprint_of( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=3): + class Local(LocalDimensionIndex): ... + """ + ) + assert len({base, swapped, counted}) == 3 + assert fingerprinting.strict_fingerprinter(V2E) != fingerprinting.strict_fingerprinter(E2V) + def test_grid_type_deduction(self): assert ( transform_utils._deduce_grid_type(None, [Vertex, V2E]) is common.GridType.UNSTRUCTURED From 17c1f757576d6e54cf86f0db9208d046943c83cf Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 00:47:19 +0200 Subject: [PATCH 03/13] feat[next]: connectivities that share another connectivity's local dimension Flattened sparse patterns (ICON4Py's C2CE, E2ECV, ...) index the same neighbor axis as another connectivity. A declaration can now adopt an owned local dimension; it is then named in the IR by its own tag (offset_tag), since the local dimension's tag already names the owner's table. --- .../ADRs/next/0029-Connectivities_As_Types.md | 33 ++++++++---- src/gt4py/next/common.py | 53 +++++++++++++------ .../test_neighbor_connectivity.py | 47 ++++++++++++++-- .../unit_tests/test_neighbor_connectivity.py | 18 ++++++- 4 files changed, 120 insertions(+), 31 deletions(-) diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index 45898fd8d7..a8d369ff7a 100644 --- a/docs/development/ADRs/next/0029-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -52,7 +52,12 @@ and skip values were never checked against the `FieldOffset` declaration. dimension can have at most one owner. - A local dimension with no table, such as the coefficient axis of a fixed-size stencil, is declared on its own: `class LsqCoeff(LocalDimensionIndex, size=3)`. - Its `owner` is `None`. + Its `owner` is `None`. A declaration can also *adopt* such a module-level local + dimension (`Local = LsqCoeff`), which then keeps its own tag. +- A connectivity can *share* another one's local dimension, `Local = C2E.Local`. + This is the flattened sparse pattern, e.g. cell-to-cell-edge (`C2CE: Cell -> CellEdge`) indexing the same neighbor axis as `C2E`, so that its results + combine with `C2E`-shaped sparse fields. The owner stays `C2E`, and the + neighbor counts and skip-value structure are the owner's. - `max_neighbors` and `min_neighbors` are optional class keywords, not type parameters: Python has no integer type parameters, and nothing static needs the count. A declared count is a constraint on the bound table; an undeclared @@ -99,15 +104,23 @@ dimension generically uses a `TypeVar` bound to `LocalDimensionIndex`. A declaration is typed like the `FieldOffset` it replaces: `V2E.__gt_type__()` is the `ts.OffsetType` of the derived offset `(Codomain -> (Origin, Local))`, -whose tag is **the local dimension's tag**, `V2E.Local.tag`. This is the single -string that shifts, neighbor reductions and sparse arguments already use to find -the table in the offset provider, so existing backends need no change. -`V2E.Local` inside DSL code types as that local dimension, and -`FieldOffset.Local` names the same thing on a legacy offset, so the spelling -works for both. The other frontend touch points treat the class like the -`FieldOffset` it derives: grid-type deduction (`transform_utils`, `past_to_itir`) -counts it as unstructured, and embedded `premap` accepts it. `V2E[i]` subscripts the -metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. +whose tag is the connectivity's `offset_tag`: + +- **the local dimension's tag**, `V2E.Local.tag`, for the connectivity that + declares it. This is the single string that shifts, neighbor reductions and + sparse arguments already use to find the table in the offset provider, so + existing backends need no change. +- **its own tag**, `C2CE.tag`, for a connectivity that shares another one's local + dimension, since the local dimension's tag already names the owner's table. + Shifts find the table by that tag; reductions and sparse arguments still find + the owner's table by the local dimension's tag, for its neighbor structure, so + the owner has to be bound too. + `V2E.Local` inside DSL code types as that local dimension, and + `FieldOffset.Local` names the same thing on a legacy offset, so the spelling + works for both. The other frontend touch points treat the class like the + `FieldOffset` it derives: grid-type deduction (`transform_utils`, `past_to_itir`) + counts it as unstructured, and embedded `premap` accepts it. `V2E[i]` subscripts the + metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. ## Consequences diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 341b11160a..bb032d8bac 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -2041,6 +2041,28 @@ def tag(cls) -> Tag: """The connectivity's identity: its qualified Python name.""" return f"{cls.__module__}.{cls.__qualname__}" + @property + def offset_tag(cls) -> Tag: + """ + The name of the connectivity in the IR, and its key in a normalized offset provider. + + The tag of its local dimension, if it declares it: shifts, neighbor reductions and sparse + arguments then all find the table under one string. A connectivity that *shares* another + one's local dimension (a flattened sparse pattern, e.g. cell-to-cell-edge indexing the + same neighbor axis as cell-to-edge) is named by its own tag, since the local dimension's + tag already names its owner's table. + """ + local = cls._local() + return local.tag if local.owner is cls else cls.tag + + def _local(cls) -> type[LocalDimensionIndex]: + if "Local" not in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}' is not a connectivity declaration; declare one by" + " subclassing 'NeighborConnectivity[Origin, Codomain]'." + ) + return cls.Local + def __call__(cls, *args: Any, **kwargs: Any) -> NoReturn: raise TypeError( f"'{cls.__qualname__}' is a connectivity declaration and cannot be instantiated;" @@ -2076,21 +2098,13 @@ def __gt_type__(cls) -> Any: def __gt_field_offset__(cls) -> Any: """ - The `FieldOffset` equivalent to this connectivity. - - Its tag is the *local dimension's* tag, the single string that shifts, neighbor - reductions and sparse arguments all use to find the table in the offset provider. + The `FieldOffset` equivalent to this connectivity, tagged with `offset_tag`. """ from gt4py.next.ffront import fbuiltins - if "Local" not in cls.__dict__: - raise TypeError( - f"'{cls.__qualname__}' is not a connectivity declaration; declare one by" - " subclassing 'NeighborConnectivity[Origin, Codomain]'." - ) if (field_offset := cls.__dict__.get("_field_offset")) is None: field_offset = fbuiltins.FieldOffset( - cls.Local.tag, source=cls.codomain, target=(cls.origin, cls.Local) + cls.offset_tag, source=cls.codomain, target=(cls.origin, cls._local()) ) type.__setattr__(cls, "_field_offset", field_offset) return field_offset @@ -2160,14 +2174,21 @@ def __init_subclass__( f"'{name}' must declare its local dimension as a nested class:" " 'class Local(LocalDimensionIndex): ...'." ) - if local.owner is not None: - raise TypeError( - f"'{name}': '{local.__qualname__}' is already the local dimension of" - f" '{local.owner.__qualname__}'." - ) - max_neighbors = _check_neighbor_count(cls, "max_neighbors", max_neighbors) min_neighbors = _check_neighbor_count(cls, "min_neighbors", min_neighbors) + if local.owner is not None: + # Sharing another connectivity's local dimension: the neighbor structure is the + # owner's, including its counts, and the sharing connectivity is named by its own tag. + if (max_neighbors, min_neighbors) != (None, None) and ( + max_neighbors, + min_neighbors, + ) != (local.max_neighbors, local.min_neighbors): + raise TypeError( + f"'{name}' shares the local dimension of '{local.owner.__qualname__}', whose" + " neighbor counts are declared by its owner." + ) + cls.origin, cls.codomain = origin, codomain + return for count_name, count in ( ("max_neighbors", max_neighbors), ("min_neighbors", min_neighbors), diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py index f8167f2627..68fd18f439 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py @@ -31,6 +31,11 @@ class V2E(gtx.NeighborConnectivity[V, E], max_neighbors=4, min_neighbors=4): class Local(gtx.LocalDimensionIndex): ... +#: A second connectivity over the same neighbor axis, bound to a different table. +class V2EShared(gtx.NeighborConnectivity[V, E]): + Local = V2E.Local + + @pytest.fixture def case(exec_alloc_descriptor): mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) @@ -43,6 +48,14 @@ def case(exec_alloc_descriptor): allocator=exec_alloc_descriptor.allocator, ) common.check_neighbor_table(V2E, table) + shared_table = constructors.as_connectivity( + domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, + codomain=E, + data=np.ascontiguousarray(v2e_arr[:, ::-1]), + skip_value=None, + allocator=exec_alloc_descriptor.allocator, + ) + common.check_neighbor_table(V2EShared, shared_table) return cases.Case( ( None @@ -51,15 +64,15 @@ def case(exec_alloc_descriptor): ), # NOTE: still keyed on the local dimension's tag; class keys come with the removal of # `FieldOffset`. - offset_provider={V2E.Local.tag: table}, + offset_provider={V2E.offset_tag: table, V2EShared.offset_tag: shared_table}, default_sizes={V: mesh.num_vertices, E: mesh.num_edges, V2E.Local: v2e_arr.shape[1]}, grid_type=common.GridType.UNSTRUCTURED, allocator=exec_alloc_descriptor.allocator, ) -def _table(case: cases.Case) -> np.ndarray: - return case.offset_provider[V2E.Local.tag].asnumpy() +def _table(case: cases.Case, connectivity=V2E) -> np.ndarray: + return case.offset_provider[connectivity.offset_tag].asnumpy() @pytest.mark.uses_unstructured_shift @@ -102,3 +115,31 @@ def testee(a: Field[Dims[E], float], out: Field[Dims[V], float]): shift_by_one(a, out=out) cases.verify_with_default_data(case, testee, lambda a: a[_table(case)[:, 0]]) + + +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_offset_tag_differing_from_local_dim +def test_shift_through_a_shared_local_dimension(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2EShared[1]) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case, V2EShared)[:, 1]]) + + +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_offset_tag_differing_from_local_dim +@pytest.mark.uses_offset_tag_differing_from_local_dim_in_reduction +def test_reduction_through_a_shared_local_dimension(case): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + # combines a sparse field on the shared axis with each connectivity's neighbors + return neighbor_sum(s * a(V2EShared) - a(V2E), axis=V2E.Local) + + cases.verify_with_default_data( + case, + testee, + lambda s, a: np.sum(s * a[_table(case, V2EShared)] - a[_table(case)], axis=1), + ) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index f0705332cd..9ff393fc1b 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -141,10 +141,10 @@ class Local(DimensionIndex, kind=DimensionKind.LOCAL): ... ), ( """ - class C(NeighborConnectivity[Vertex, Edge]): + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=5): Local = V2E.Local """, - "already the local dimension of 'V2E'", + "counts are declared by its owner", ), ( """ @@ -227,6 +227,20 @@ class C(NeighborConnectivity[Vertex, Edge]): assert ns["Coeff"].owner is ns["C"] assert (ns["Coeff"].max_neighbors, ns["Coeff"].min_neighbors) == (3, 3) + def test_sharing_a_local_dimension(self): + ns = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local = V2E.Local + """ + ) + shared = ns["V2EShared"] + assert shared.Local is V2E.Local + assert V2E.Local.owner is V2E + assert V2E.offset_tag == V2E.Local.tag + assert shared.offset_tag == shared.tag + assert shared.__gt_type__().tag == shared.tag + def test_non_integer_index(self): with pytest.raises(TypeError, match="indexed by an integer"): V2E[Vertex] From 8d6766fac0a27e30e0d5ebd5f95df7dfdc826841 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 15:42:57 +0200 Subject: [PATCH 04/13] fix[next]: a declaration's Local stays a type for both checkers An annotated 'Local' -- on 'NeighborConnectivity' or on its metaclass -- makes every declaration's local dimension a *variable*: pyright then rejects 'Field[Dims[V, V2E.Local], float]', and mypy rejects an adopted or shared one. Annotate it nowhere; library code reads it through 'common.local_dimension_of', and adoption and sharing are written 'Local: TypeAlias = ...'. Also from the review: a redefined declaration re-owns an adopted local dimension instead of becoming a sharer, and its counts are checked against the local dimension's own 'size='. The missing-'Local' error names both spellings. Adds pyright over 'typing_tests/pyright_probes.py' to the typing session, which is what catches a regression here: the mypy cases cannot. --- .../ADRs/next/0029-Connectivities_As_Types.md | 36 ++++-- noxfile.py | 3 + pyproject.toml | 1 + src/gt4py/next/__init__.py | 2 + src/gt4py/next/common.py | 61 +++++++--- src/gt4py/next/fingerprinting.py | 11 +- .../unit_tests/test_neighbor_connectivity.py | 53 ++++++++- typing_tests/pyright_probes.py | 111 ++++++++++++++++++ typing_tests/pyrightconfig.json | 5 + typing_tests/test_next.yaml | 6 +- uv.lock | 15 +++ 11 files changed, 265 insertions(+), 39 deletions(-) create mode 100644 typing_tests/pyright_probes.py create mode 100644 typing_tests/pyrightconfig.json diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index a8d369ff7a..150976307b 100644 --- a/docs/development/ADRs/next/0029-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -49,12 +49,15 @@ and skip values were never checked against the `FieldOffset` declaration. - The local dimension is the nested class `Local`, a subclass of `LocalDimensionIndex`. Declaring it is required, and `NeighborConnectivity` sets `Local.owner` to the connectivity when the class is created. A local - dimension can have at most one owner. + dimension can have at most one owner; a declaration redefined under the same + name (a re-run notebook cell) takes ownership over again, and for a local + dimension adopted rather than nested, the first declaration wins. - A local dimension with no table, such as the coefficient axis of a fixed-size stencil, is declared on its own: `class LsqCoeff(LocalDimensionIndex, size=3)`. Its `owner` is `None`. A declaration can also *adopt* such a module-level local - dimension (`Local = LsqCoeff`), which then keeps its own tag. -- A connectivity can *share* another one's local dimension, `Local = C2E.Local`. + dimension, written `Local: TypeAlias = LsqCoeff`, which then keeps its own tag. +- A connectivity can *share* another one's local dimension, + `Local: TypeAlias = C2E.Local`. This is the flattened sparse pattern, e.g. cell-to-cell-edge (`C2CE: Cell -> CellEdge`) indexing the same neighbor axis as `C2E`, so that its results combine with `C2E`-shaped sparse fields. The owner stays `C2E`, and the neighbor counts and skip-value structure are the owner's. @@ -90,15 +93,24 @@ tree already distinguishes local dimensions by a runtime `kind` check, so it keeps doing so; generic constructors whose parameter must be a primary dimension (`NeighborConnectivity[Origin, Codomain]`, `Staggered[D]`) check it at runtime. -### How the base reaches `Local` - -The base class declares `Local: ClassVar[type[LocalDimensionIndex]]` as an -annotation only. A real nested class on the base would be an incompatible -override for pyright in every declaration. The annotation makes `conn.Local` a -value of type `type[LocalDimensionIndex]` for library code taking any -connectivity; it does not make `conn.Local` usable as an *annotation* when -`conn` is generic, which no checker allows. Code that needs to name a local -dimension generically uses a `TypeVar` bound to `LocalDimensionIndex`. +### `Local` is not annotated anywhere + +Neither `NeighborConnectivity` nor `ConnectivityMeta` annotates `Local`, and +that is load-bearing: an annotation makes a declaration's `Local` a *variable* +for the checkers, so `Field[Dims[Vertex, V2E.Local], float]` is rejected by +pyright ("Variable not allowed in type expression") for a nested `Local`, and by +mypy ("not valid as a type") for an adopted or shared one. A real nested `Local` +on the base is not an option either: pyright reports an incompatible override in +every declaration. With no annotation, all three spellings are types for both +checkers, which `typing_tests/pyright_probes.py` pins for pyright and +`typing_tests/test_next.yaml` for mypy. + +The cost is that `conn.Local` is not an attribute the checkers know for a +*generic* `conn`. Library code reads it through `common.local_dimension_of(conn)` +instead, and code that has to name a local dimension generically uses a +`TypeVar` bound to `LocalDimensionIndex`. Writing an adopted or shared local as +`Local: TypeAlias = ...` (rather than a plain assignment) is what keeps mypy +treating it as a type. ### Frontend integration diff --git a/noxfile.py b/noxfile.py index c8091de09e..31327f076d 100755 --- a/noxfile.py +++ b/noxfile.py @@ -349,6 +349,9 @@ def test_typing_exports(session: nox.Session) -> None: "typing_tests", *session.posargs, ) + # A second checker, on code that must type-check for a downstream user: mypy and pyright + # disagree about what counts as a type, which the mypy-only cases above cannot catch. + session.run("pyright", "--project", "typing_tests", "typing_tests/pyright_probes.py") # -- DaCe codegen determinism check -- diff --git a/pyproject.toml b/pyproject.toml index e131ebda0e..9072d8c5df 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -65,6 +65,7 @@ typing = [ typing_exports = [ # to test typing with gt4py in downstream code {include-group = "typing"}, + 'pyright>=1.1.400', # the second checker: it disagrees with mypy about what counts as a type 'types-six>=1.17.0.20251009', # can not let mypy auto-install types as that leads to unexpected stderr output (which means test failure) 'pytest-mypy-plugins>=4.0.0', # pytest plugin for running mypy on code snippets "xarray>=2024.1.0" # one of the regression tests requires xarray diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index c41ce38680..253e574f95 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -41,6 +41,7 @@ domain, flip_staggered, is_staggered, + local_dimension_of, resolve, unit_range, ) @@ -140,6 +141,7 @@ "unit_range", "UnitRange", "is_staggered", + "local_dimension_of", "flip_staggered", "as_non_staggered", # from constructors diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index bb032d8bac..a29b70a1b0 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1994,6 +1994,8 @@ class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): #: Least number of *valid* neighbors of any element, if declared. Fewer than #: `max_neighbors` means the table pads with skip values. min_neighbors: ClassVar[Optional[int]] = None + #: The `size=` of this declaration, kept apart from the counts an owner writes below. + declared_size: ClassVar[Optional[int]] = None def __init_subclass__( cls, @@ -2011,7 +2013,8 @@ def __init_subclass__( # NOTE: reset rather than inherited: a subclass of an owned local dimension is a # different dimension, and does not index its parent's table. cls.owner = None - cls.max_neighbors = cls.min_neighbors = _check_neighbor_count(cls, "size", size) + cls.declared_size = _check_neighbor_count(cls, "size", size) + cls.max_neighbors = cls.min_neighbors = cls.declared_size def _check_neighbor_count(cls: type, name: str, count: Optional[int]) -> Optional[int]: @@ -2032,7 +2035,10 @@ class ConnectivityMeta(type): (`a(V2E)`, `a(V2E[0])`), and it is what the neighbor table bound at call time must match. """ - Local: type[LocalDimensionIndex] + # NOTE: `Local` is deliberately *not* annotated here, nor on `NeighborConnectivity`: an + # annotated `Local` makes every declaration's nested class a *variable* for the checkers, so + # `Field[Dims[V, V2E.Local]]` is rejected (pyright) or "not valid as a type" (mypy, for the + # assigned form). Library code reads it through `local_dimension_of`. origin: Dimension codomain: Dimension @@ -2056,12 +2062,12 @@ def offset_tag(cls) -> Tag: return local.tag if local.owner is cls else cls.tag def _local(cls) -> type[LocalDimensionIndex]: - if "Local" not in cls.__dict__: + if (local := cls.__dict__.get("Local")) is None: raise TypeError( f"'{cls.__qualname__}' is not a connectivity declaration; declare one by" " subclassing 'NeighborConnectivity[Origin, Codomain]'." ) - return cls.Local + return cast(type[LocalDimensionIndex], local) def __call__(cls, *args: Any, **kwargs: Any) -> NoReturn: raise TypeError( @@ -2132,9 +2138,8 @@ class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( (True, 6, 5) """ - # NOTE: an annotation only, never assigned here. A real nested class on the base would be - # flagged by pyright as an incompatible override in every declaration. - Local: ClassVar[type[LocalDimensionIndex]] + # NOTE: `Local` is not annotated (see `ConnectivityMeta`); every subclass declares it, as a + # nested class or as `Local: TypeAlias = `. origin: ClassVar[Dimension] codomain: ClassVar[Dimension] @@ -2171,12 +2176,17 @@ def __init_subclass__( local = cls.__dict__.get("Local") if not (isinstance(local, DimensionMeta) and issubclass(local, LocalDimensionIndex)): raise TypeError( - f"'{name}' must declare its local dimension as a nested class:" - " 'class Local(LocalDimensionIndex): ...'." + f"'{name}' must declare its local dimension, either as a nested class" + " ('class Local(LocalDimensionIndex): ...') or by adopting one" + " ('Local: TypeAlias = SomeLocalDim')." ) max_neighbors = _check_neighbor_count(cls, "max_neighbors", max_neighbors) min_neighbors = _check_neighbor_count(cls, "min_neighbors", min_neighbors) - if local.owner is not None: + # NOTE: a declaration whose tag is the owner's is a *redefinition* of it (a re-run + # notebook cell), not a second connectivity sharing the local dimension, so it takes + # ownership over again. Ownership of a local dimension that is adopted, rather than + # nested, otherwise goes to whoever declares first. + if local.owner is not None and local.owner.tag != cls.tag: # Sharing another connectivity's local dimension: the neighbor structure is the # owner's, including its counts, and the sharing connectivity is named by its own tag. if (max_neighbors, min_neighbors) != (None, None) and ( @@ -2193,14 +2203,19 @@ def __init_subclass__( ("max_neighbors", max_neighbors), ("min_neighbors", min_neighbors), ): - declared = getattr(local, count_name) - if count is not None and declared is not None and count != declared: + # NOTE: against the local dimension's own `size=`, not against counts a previous + # owner wrote: a redefinition must be checked against what its `Local` declares. + if ( + count is not None + and local.declared_size is not None + and count != local.declared_size + ): raise TypeError( f"'{name}': '{count_name}={count}' contradicts the size declared by" - f" '{local.__qualname__}' ({declared})." + f" '{local.__qualname__}' ({local.declared_size})." ) - max_neighbors = max_neighbors if max_neighbors is not None else local.max_neighbors - min_neighbors = min_neighbors if min_neighbors is not None else local.min_neighbors + max_neighbors = max_neighbors if max_neighbors is not None else local.declared_size + min_neighbors = min_neighbors if min_neighbors is not None else local.declared_size if ( max_neighbors is not None and min_neighbors is not None @@ -2216,6 +2231,20 @@ def __init_subclass__( local.max_neighbors, local.min_neighbors = max_neighbors, min_neighbors +def local_dimension_of(connectivity: type[NeighborConnectivity]) -> type[LocalDimensionIndex]: + """ + The local dimension a connectivity declares, adopts or shares. + + Library code reads `V2E.Local` through this accessor: the attribute is intentionally not + annotated, so that a declaration's `Local` stays a *type* for the type checkers (see + `ConnectivityMeta`). + + Raises: + TypeError: If `connectivity` declares no local dimension. + """ + return cast(ConnectivityMeta, connectivity)._local() + + def check_neighbor_table( connectivity: type[NeighborConnectivity], table: NeighborTable | NeighborConnectivityType, @@ -2235,7 +2264,7 @@ def check_neighbor_table( """ table_type = table if isinstance(table, NeighborConnectivityType) else table.__gt_type__() name = connectivity.__qualname__ - local = connectivity.Local + local = local_dimension_of(connectivity) def fail(reason: str) -> NoReturn: raise ValueError(f"The table bound to '{name}' does not match its declaration: {reason}.") diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index a9d8e206b3..ce78404868 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -235,14 +235,15 @@ def _instance_state(obj: Any) -> tuple[Any, ...]: # by what it declares: redefining `V2E` under the same name with other dimensions or counts # (e.g. re-running a notebook cell) must not reuse artifacts built for the old declaration. common.ConnectivityMeta: lambda obj: ( - # NOTE: the name goes into the state unverified; importability is still enforced by the - # strict fingerprinter through `obj.Local`, which is a class nested in `obj`. + # NOTE: the name goes into the state unverified; the local dimension is fingerprinted as a + # class, so the strict fingerprinter still checks *its* importability (which is the + # connectivity's own, for a nested `Local`, and another module's for a shared one). Deconstruction.from_pieces( obj.origin, obj.codomain, - obj.Local, - obj.Local.max_neighbors, - obj.Local.min_neighbors, + common.local_dimension_of(obj), + common.local_dimension_of(obj).max_neighbors, + common.local_dimension_of(obj).min_neighbors, state=b"neighbor_connectivity\0" + obj.tag.encode(), ) if "Local" in obj.__dict__ diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 9ff393fc1b..5d09dd8d76 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -55,6 +55,7 @@ def _declare(source: str) -> dict: """ namespace = { "__name__": __name__, + "typing": typing, "DimensionIndex": DimensionIndex, "DimensionKind": DimensionKind, "LocalDimensionIndex": LocalDimensionIndex, @@ -63,6 +64,7 @@ def _declare(source: str) -> dict: "Edge": Edge, "KDim": KDim, "V2E": V2E, + "ConstListDim": common.ConstListDim, } exec(textwrap.dedent(source), namespace) return namespace @@ -142,7 +144,7 @@ class Local(DimensionIndex, kind=DimensionKind.LOCAL): ... ( """ class C(NeighborConnectivity[Vertex, Edge], max_neighbors=5): - Local = V2E.Local + Local: typing.TypeAlias = V2E.Local """, "counts are declared by its owner", ), @@ -221,7 +223,7 @@ def test_adopting_an_ownerless_local(self): class Coeff(LocalDimensionIndex, size=3): ... class C(NeighborConnectivity[Vertex, Edge]): - Local = Coeff + Local: typing.TypeAlias = Coeff """ ) assert ns["Coeff"].owner is ns["C"] @@ -231,7 +233,7 @@ def test_sharing_a_local_dimension(self): ns = _declare( """ class V2EShared(NeighborConnectivity[Vertex, Edge]): - Local = V2E.Local + Local: typing.TypeAlias = V2E.Local """ ) shared = ns["V2EShared"] @@ -253,7 +255,7 @@ def test_function_local_declaration(self): with pytest.raises(TypeError, match="module level"): class C(NeighborConnectivity[Vertex, Edge]): - Local = LsqCoeff + Local: typing.TypeAlias = LsqCoeff def test_subclass_of_owned_local_is_ownerless(self): ns = _declare( @@ -428,3 +430,46 @@ def test_grid_type_deduction(self): ) with pytest.raises(ValueError, match="CARTESIAN"): transform_utils._deduce_grid_type(common.GridType.CARTESIAN, [V2E]) + + +def test_redefined_declaration_with_an_adopted_local(monkeypatch): + """Re-running a cell must re-own the adopted local dimension, not become a sharer.""" + import sys + import types as pytypes + + module = pytypes.ModuleType("_readopted_connectivity_module") + monkeypatch.setitem(sys.modules, module.__name__, module) + source = textwrap.dedent( + """ + import typing + + from gt4py.next.common import DimensionIndex, LocalDimensionIndex, NeighborConnectivity + + class V(DimensionIndex): ... + class E(DimensionIndex): ... + class V2EDim(LocalDimensionIndex, size={n}): ... + class V2E(NeighborConnectivity[V, E], max_neighbors={n}): + Local: typing.TypeAlias = V2EDim + """ + ) + exec(source.format(n=4), module.__dict__) + assert module.V2E.offset_tag == module.V2EDim.tag + + # the redefinition takes ownership over again, and its counts are checked against `size=` + exec(source.format(n=2), module.__dict__) + assert module.V2EDim.owner is module.V2E + assert module.V2E.offset_tag == module.V2EDim.tag + assert module.V2EDim.max_neighbors == 2 + + +def test_local_dimension_of(): + shared = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2EShared"] + assert common.local_dimension_of(V2E) is V2E.Local + assert common.local_dimension_of(shared) is V2E.Local + with pytest.raises(TypeError, match="not a connectivity declaration"): + common.local_dimension_of(NeighborConnectivity) diff --git a/typing_tests/pyright_probes.py b/typing_tests/pyright_probes.py new file mode 100644 index 0000000000..e9be2423d9 --- /dev/null +++ b/typing_tests/pyright_probes.py @@ -0,0 +1,111 @@ +# 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 + +""" +Client code that has to type-check under *pyright*, checked by `nox -s test_typing_exports`. + +The cases in `test_next.yaml` run under mypy only, and the two checkers disagree about what +counts as a type: an annotated `Local` on a connectivity or its metaclass makes every +declaration's local dimension a *variable* for pyright, so `Field[Dims[V, V2E.Local], float]` +is rejected there while mypy accepts it (see ADR 0029). Everything here must be error-free. +""" + +from __future__ import annotations + +import typing + +from gt4py import next as gtx + + +class Vertex(gtx.DimensionIndex): ... + + +class Edge(gtx.DimensionIndex): ... + + +class Cell(gtx.DimensionIndex): ... + + +class CellEdge(gtx.DimensionIndex): ... + + +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): + class Local(gtx.LocalDimensionIndex): ... + + +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + +#: A flattened sparse pattern sharing `C2E`'s neighbor axis. +class C2CE(gtx.NeighborConnectivity[Cell, CellEdge]): + Local: typing.TypeAlias = C2E.Local + + +class LsqCoeff(gtx.LocalDimensionIndex, size=3): ... + + +#: A declaration adopting a local dimension declared at module level. +class V2EAdopted(gtx.NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = LsqCoeff + + +def nested_local(sparse: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + +def shared_local(sparse: gtx.Field[gtx.Dims[Cell, C2CE.Local], gtx.float64]) -> None: ... + + +def adopted_local(sparse: gtx.Field[gtx.Dims[Vertex, V2EAdopted.Local], gtx.float64]) -> None: ... + + +def a_shared_local_is_its_owners( + owned: gtx.Field[gtx.Dims[Cell, C2E.Local], gtx.float64], + shared: gtx.Field[gtx.Dims[Cell, C2CE.Local], gtx.float64], +) -> None: + shared_local(owned) # the two spellings are one type + shared_local(shared) + + +def an_adopted_local_is_the_adopted_one( + coefficients: gtx.Field[gtx.Dims[Vertex, LsqCoeff], gtx.float64], +) -> None: + adopted_local(coefficients) + + +L = typing.TypeVar("L", bound=gtx.LocalDimensionIndex) + + +def local_of(connectivity: type[gtx.NeighborConnectivity]) -> type[gtx.LocalDimensionIndex]: + # generic code names a local dimension through the accessor, not through `conn.Local` + return gtx.local_dimension_of(connectivity) + + +def first_neighbor(sparse: gtx.Field[gtx.Dims[Vertex, L], gtx.float64]) -> type[L]: + raise NotImplementedError + + +def generic_local(sparse: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: + typing.assert_type(first_neighbor(sparse), type[V2E.Local]) + + +@gtx.field_operator +def reduce_over_a_local_dimension( + a: gtx.Field[gtx.Dims[Edge], gtx.float64], +) -> gtx.Field[gtx.Dims[Vertex], gtx.float64]: + return gtx.neighbor_sum(a(V2E), axis=V2E.Local) + + +@gtx.field_operator +def shift_by_a_dimension( + a: gtx.Field[gtx.Dims[KDim], gtx.float64], +) -> gtx.Field[gtx.Dims[KDim], gtx.float64]: + return a(KDim + 1) diff --git a/typing_tests/pyrightconfig.json b/typing_tests/pyrightconfig.json new file mode 100644 index 0000000000..4f44d9cfb8 --- /dev/null +++ b/typing_tests/pyrightconfig.json @@ -0,0 +1,5 @@ +{ + "typeCheckingMode": "standard", + "reportMissingImports": "error", + "reportMissingTypeStubs": "none" +} diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 04007c4a32..a65d3e7d09 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -362,7 +362,9 @@ L = typing.TypeVar("L", bound=gtx.LocalDimensionIndex) def local_of(conn: type[gtx.NeighborConnectivity]) -> type[gtx.LocalDimensionIndex]: - return conn.Local + # `conn.Local` is not annotated, so that a declaration's `Local` stays a type; generic + # code reads it through the accessor + return gtx.local_dimension_of(conn) def first(a: gtx.Field[gtx.Dims[Vertex, L], gtx.float64]) -> type[L]: raise NotImplementedError @@ -370,4 +372,4 @@ def caller(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: reveal_type(first(a)) out: | - main:20:17: note: Revealed type is "type[main.V2E.Local]" + main:22:17: note: Revealed type is "type[main.V2E.Local]" diff --git a/uv.lock b/uv.lock index ef2d5ca914..648636a40e 100644 --- a/uv.lock +++ b/uv.lock @@ -1442,6 +1442,7 @@ typing = [ ] typing-exports = [ { name = "mypy", extra = ["faster-cache"] }, + { name = "pyright" }, { name = "pytest-mypy-plugins" }, { name = "types-decorator" }, { name = "types-docutils" }, @@ -1598,6 +1599,7 @@ typing = [ ] typing-exports = [ { name = "mypy", extras = ["faster-cache"], specifier = ">=1.13.0" }, + { name = "pyright", specifier = ">=1.1.400" }, { name = "pytest-mypy-plugins", specifier = ">=4.0.0" }, { name = "types-decorator", specifier = ">=5.1.8" }, { name = "types-docutils", specifier = ">=0.21.0" }, @@ -3126,6 +3128,19 @@ version = "2.1" source = { registry = "https://pypi.org/simple" } sdist = { url = "https://files.pythonhosted.org/packages/bc/7c/d724ef1ec3ab2125f38a1d53285745445ec4a8f19b9bb0761b4064316679/pyreadline-2.1.zip", hash = "sha256:4530592fc2e85b25b1a9f79664433da09237c1a270e4d78ea5aa3a2c7229e2d1", size = 109189, upload-time = "2015-09-16T08:24:48.745Z" } +[[package]] +name = "pyright" +version = "1.1.414" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nodeenv" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/e1/1b/244c7b710031ada80f27e579ec20d28a2285dfc318fed0339866b1047f12/pyright-1.1.414.tar.gz", hash = "sha256:523c0a97c60da6333234955c277730c9cf4f5bd6d5399e7b7d2b0fc5d3599524", size = 4154638, upload-time = "2026-09-10T12:26:53.181Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d7/ba/18b6e682ead424ad24bcc134339ae5d1b931cd9ae260540592a058a91279/pyright-1.1.414-py3-none-any.whl", hash = "sha256:2a6b4b3298c9eec174c5ed83bd338de6eee82df2992f3e1930e6199d381be36f", size = 6225049, upload-time = "2026-09-10T12:26:51.427Z" }, +] + [[package]] name = "pytest" version = "9.1.1" From cff9803f2b5824abce43ccb91b09bb4f1c621daf Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 16:19:56 +0200 Subject: [PATCH 05/13] docs[next]: the neighbor-count TODO is partly done by the declaration --- src/gt4py/next/common.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index a29b70a1b0..f579fff507 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1302,7 +1302,10 @@ def has_skip_values(self) -> bool: @dataclasses.dataclass(frozen=True) class NeighborConnectivityType(ConnectivityType): - # TODO(havogt): refactor towards encoding this information in the local dimensions of the ConnectivityType.domain + # NOTE: partly encoded in the local dimension since ADR 0029: a `LocalDimensionIndex` carries + # `max_neighbors` / `min_neighbors` where the declaration states them, and this record is + # checked against them (`check_neighbor_table`). It stays the *bound* count, which a + # declaration may leave to the table. max_neighbors: int @property From 9dcf55c22f88b290f5bf06029f6b7932e4e201ab Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 24 Sep 2026 00:03:11 +0200 Subject: [PATCH 06/13] refactor[next]: ConstList is an owner-less local dimension of size 1 Now that local dimensions can state their size, the local dimension of 'make_const_list' results is one; a declaration cannot adopt it, since it belongs to no connectivity. --- src/gt4py/next/common.py | 34 +++++++++++-------- .../unit_tests/test_neighbor_connectivity.py | 12 ++++++- 2 files changed, 31 insertions(+), 15 deletions(-) diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index f579fff507..1bfd4a1747 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1886,20 +1886,6 @@ def __init_subclass__(cls, /, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) -class ConstList(DimensionIndex, kind=DimensionKind.LOCAL): - """ - The local dimension of a list of one repeated value (`make_const_list`). - - The value is broadcast against the neighbor lists it is combined with, and a materialized - constant list has extent 1 along it. It indexes no table, so it is never in an offset provider. - - Declared here, once: it used to be built independently in `iterator/embedded.py` and in the - DaCe lowering, which only worked while dimensions compared by `(name, kind)`. - """ - - __slots__ = () - - def _reduce_staggered(cls: StaggeredMeta) -> Any: """ Pickle a staggered dimension through its base, falling back to by-reference. @@ -2030,6 +2016,21 @@ def _check_neighbor_count(cls: type, name: str, count: Optional[int]) -> Optiona return int(count) +class ConstList(LocalDimensionIndex, size=1): + """ + The local dimension of a list of one repeated value (`make_const_list`). + + An owner-less local dimension of size 1: the value is broadcast against the neighbor lists it + is combined with, and a materialized constant list has extent 1 along it. It indexes no table, + so it is never in an offset provider. + + Declared here, once: it used to be built independently in `iterator/embedded.py` and in the + DaCe lowering, which only worked while dimensions compared by `(name, kind)`. + """ + + __slots__ = () + + class ConnectivityMeta(type): """ Metaclass of `NeighborConnectivity` declarations. @@ -2183,6 +2184,11 @@ def __init_subclass__( " ('class Local(LocalDimensionIndex): ...') or by adopting one" " ('Local: TypeAlias = SomeLocalDim')." ) + if local is ConstList: + raise TypeError( + f"'{name}' cannot adopt '{ConstList.__qualname__}': it is the local dimension" + " of 'make_const_list' results and belongs to no connectivity." + ) max_neighbors = _check_neighbor_count(cls, "max_neighbors", max_neighbors) min_neighbors = _check_neighbor_count(cls, "min_neighbors", min_neighbors) # NOTE: a declaration whose tag is the owner's is a *redefinition* of it (a re-run diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 5d09dd8d76..c5f3a78743 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -64,7 +64,7 @@ def _declare(source: str) -> dict: "Edge": Edge, "KDim": KDim, "V2E": V2E, - "ConstListDim": common.ConstListDim, + "ConstList": common.ConstList, } exec(textwrap.dedent(source), namespace) return namespace @@ -473,3 +473,13 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): assert common.local_dimension_of(shared) is V2E.Local with pytest.raises(TypeError, match="not a connectivity declaration"): common.local_dimension_of(NeighborConnectivity) + + +def test_the_const_list_dimension_cannot_be_adopted(): + with pytest.raises(TypeError, match="cannot adopt"): + _declare( + """ + class C(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = ConstList + """ + ) From a288e6047cf14f3223bb6ea2736e1b3a1a761836 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 01:03:00 +0200 Subject: [PATCH 07/13] feat[next]: backends support connectivities sharing a local dimension Reductions, sparse arguments and list materialization used to find a connectivity by the local dimension's tag, which only names the owner's table. They now look up a table over the local dimension (common.connectivity_key_over), so a connectivity sharing another one's local dimension works on every backend, bound on its own or together with its owner. DaCe sizes a connectivity array's local dimension from its own table. Lifts the PR 1 xfail markers. --- pyproject.toml | 2 - src/gt4py/next/common.py | 32 ++++++++++++++ src/gt4py/next/embedded/nd_array_field.py | 4 +- src/gt4py/next/iterator/embedded.py | 40 ++++++++++++++---- .../next/iterator/transforms/unroll_reduce.py | 10 +++-- .../codegens/gtfn/gtfn_module.py | 7 +++- .../runners/dace/lowering/gtir_to_sdfg.py | 16 +++++-- .../dace/lowering/gtir_to_sdfg_lambda.py | 16 +++++-- .../runners/dace/sdfg_args.py | 21 ++++++++++ tests/next_tests/definitions.py | 39 +---------------- .../test_neighbor_connectivity.py | 42 +++++++++++++++++-- .../test_offset_dimensions_names.py | 24 ++++------- .../transforms_tests/test_unroll_reduce.py | 8 ++-- 13 files changed, 174 insertions(+), 87 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 9072d8c5df..ba7f6cd74d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -253,8 +253,6 @@ markers = [ 'uses_ir_if_stmts', 'uses_lift: tests that require backend support for lift builtin function', 'uses_negative_modulo: tests that require backend support for modulo on negative numbers', - 'uses_offset_tag_differing_from_local_dim: tests that shift with an offset tag that differs from the name of its local dimension', - 'uses_offset_tag_differing_from_local_dim_in_reduction: tests that reduce over an offset whose tag differs from the name of its local dimension', 'uses_origin: tests that require backend support for domain origin', 'uses_reduce_with_lambda: tests that use lambdas as reduce functions', 'uses_reduction_with_only_sparse_fields: tests that require backend support for with sparse fields', diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 1bfd4a1747..58729c4184 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1543,6 +1543,38 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid get_offset_type: Callable[[OffsetProviderType, str], OffsetProviderTypeElem] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and OffsetProviderType overlap +def connectivity_key_over( + offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension +) -> str: + """ + The key of a bound connectivity whose local dimension is `local_dim`. + + Neighbor reductions and sparse fields know only their local dimension, and use its table + for the neighbor count and the skip values. That is the table keyed by the local dimension's + tag, i.e. its owner's, if bound. Otherwise it is a connectivity *sharing* the local dimension + (see `NeighborConnectivity`), keyed by its own tag, which has the same neighbor structure. + + Raises: + KeyError: If no bound connectivity has `local_dim` as its local dimension. + """ + if local_dim.tag in offset_provider: + return local_dim.tag + for key, connectivity in offset_provider.items(): + if isinstance(connectivity, NeighborConnectivityType): + neighbor_dim = connectivity.neighbor_dim + elif is_neighbor_table(connectivity): + neighbor_dim = connectivity.domain.dims[1] + else: + continue + if neighbor_dim is local_dim: + assert isinstance(key, str) + return key + raise KeyError( + f"No connectivity over the local dimension '{local_dim.tag}' is bound in the offset" + f" provider, which has {sorted(map(str, offset_provider))}." + ) + + def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: str) -> bool: """Determine if offset provider has an element for the given offset tag.""" try: diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index 38e760995c..8a1ee0a79c 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -979,8 +979,8 @@ def _builtin_op( current_offset_provider = embedded_context.get_offset_provider(None) assert current_offset_provider is not None offset_definition = common.get_offset( - current_offset_provider, axis.tag - ) # assumes offset and local dimension have same name + current_offset_provider, common.connectivity_key_over(current_offset_provider, axis) + ) assert common.is_neighbor_table(offset_definition) new_domain = common.Domain(*[nr for nr in field.domain if nr.dim != axis]) diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index b0bbf16f63..985bc139cc 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -569,7 +569,12 @@ def execute_shift( if tag == common.ConstList.tag: new_entry[i] = 0 else: - offset_implementation = common.get_offset(offset_provider, tag) + # NOTE: the sparse tag is the local dimension's; the table over it may be + # keyed by a connectivity sharing it (see `common.connectivity_key_over`). + offset_implementation = common.get_offset( + offset_provider, + common.connectivity_key_over(offset_provider, common.resolve(tag)), + ) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().source_dim cur_index = pos[source_dim.tag] @@ -1000,7 +1005,7 @@ def field_setitem(self, named_indices: NamedFieldIndices, value: Any): if isinstance(value, _List): for i, v in enumerate(value): # type:ignore[var-annotated, arg-type] self._ndarrayfield[ - self._translate_named_indices({**named_indices, value.offset.value: i}) # type: ignore[dict-item] + self._translate_named_indices({**named_indices, value.local_dim.tag: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ @@ -1411,16 +1416,26 @@ def __getitem__(self, i: int): return self.values[i] def __gt_type__(self) -> ts.ListType: - offset_tag = self.offset.value - assert isinstance(offset_tag, str) element_type = type_translation.from_value(self.values[0]) assert isinstance(element_type, ts.DataType) + return ts.ListType(element_type=element_type, offset_type=self.local_dim) + + @property + def local_dim(self) -> common.Dimension: + """ + The local dimension the list runs along. + + The neighbor dimension of the connectivity the list was built with, which is not + necessarily named like its offset: a connectivity can share another one's local + dimension. + """ + offset_tag = self.offset.value + assert isinstance(offset_tag, str) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None connectivity = common.get_offset(offset_provider, offset_tag) assert common.is_neighbor_table(connectivity) - local_dim = connectivity.__gt_type__().neighbor_dim - return ts.ListType(element_type=element_type, offset_type=local_dim) + return connectivity.__gt_type__().neighbor_dim @dataclasses.dataclass(frozen=True) @@ -1523,7 +1538,12 @@ def deref(self) -> Any: ) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None - connectivity = common.get_offset(offset_provider, self.list_offset) + # NOTE: `list_offset` is the local dimension's tag; the table over it may be keyed by a + # connectivity sharing it (see `common.connectivity_key_over`). + connectivity_key = common.connectivity_key_over( + offset_provider, common.resolve(self.list_offset) + ) + connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( values=tuple( @@ -1533,7 +1553,7 @@ def deref(self) -> Any: shifted := self.it.shift(*self.offsets, SparseTag(self.list_offset), i) ).can_deref() ), - offset=runtime.Offset(value=self.list_offset), + offset=runtime.Offset(value=connectivity_key), ) def can_deref(self) -> bool: @@ -1780,7 +1800,9 @@ def _fieldspec_list_to_value( offset_provider = embedded_context.get_offset_provider() offset_type = type_.offset_type assert isinstance(offset_type, common.DimensionMeta) - connectivity = common.get_offset(offset_provider, offset_type.tag) + connectivity = common.get_offset( + offset_provider, common.connectivity_key_over(offset_provider, offset_type) + ) assert common.is_neighbor_table(connectivity) return domain.insert( len(domain), diff --git a/src/gt4py/next/iterator/transforms/unroll_reduce.py b/src/gt4py/next/iterator/transforms/unroll_reduce.py index cfb7bb3226..22491346bb 100644 --- a/src/gt4py/next/iterator/transforms/unroll_reduce.py +++ b/src/gt4py/next/iterator/transforms/unroll_reduce.py @@ -40,11 +40,11 @@ def _get_neighbors_args(reduce_args: Iterable[itir.Expr]) -> Iterator[itir.FunCa return filter(_is_neighbors_or_lifted_and_neighbors, flat_reduce_args) -def _get_partial_offset_tags(reduce_args: Iterable[itir.Expr]) -> Iterable[str]: +def _get_partial_local_dims(reduce_args: Iterable[itir.Expr]) -> Iterable[common.Dimension]: assert all(isinstance(arg.type, ts.ListType) for arg in reduce_args) return [ - arg.type.offset_type.tag # type: ignore[union-attr] # checked in previous lines + arg.type.offset_type # type: ignore[union-attr] # checked in previous lines for arg in reduce_args if arg.type.offset_type is not None # type: ignore[union-attr] # checked in previous lines ] @@ -59,8 +59,10 @@ def _get_connectivity( raise ValueError("Expected a call to a 'reduce' object, i.e. 'reduce(...)(...)'.") connectivities: list[common.NeighborConnectivityType] = [] - for o in _get_partial_offset_tags(applied_reduce_node.args): - conn = common.get_offset_type(offset_provider_type, o) + for local_dim in _get_partial_local_dims(applied_reduce_node.args): + conn = common.get_offset_type( + offset_provider_type, common.connectivity_key_over(offset_provider_type, local_dim) + ) assert isinstance(conn, common.NeighborConnectivityType) connectivities.append(conn) diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index 68591309c6..78f4c7190c 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -92,7 +92,12 @@ def _process_regular_arguments( # NOTE: the tag is the offset-provider key, and its mangled form names the # `generated::_t` tag type. A legacy `FieldOffset` carries it as `value`. dim_name = dim.value if isinstance(dim, fbuiltins.FieldOffset) else dim.tag - connectivity = common.get_offset_type(offset_provider_type, dim_name) + connectivity = common.get_offset_type( + offset_provider_type, + dim_name + if isinstance(dim, fbuiltins.FieldOffset) + else common.connectivity_key_over(offset_provider_type, dim), + ) assert isinstance(connectivity, common.NeighborConnectivityType) size = connectivity.max_neighbors arg = f"gridtools::sid::dimension_to_tuple_like({arg})" diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index c4b1526bdc..d3026b3ea2 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -89,6 +89,11 @@ class DataflowBuilder(Protocol): @abc.abstractmethod def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: ... + @abc.abstractmethod + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + """The offset of a connectivity over `local_dim`, see `common.connectivity_key_over`.""" + ... + @abc.abstractmethod def unique_nsdfg_name(self, prefix: str) -> str: ... @@ -564,6 +569,9 @@ class GTIRToSDFG(eve.NodeVisitor, SDFGBuilder): def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: return gtx_common.get_offset_type(self.offset_provider_type, offset) + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + return gtx_common.connectivity_key_over(self.offset_provider_type, local_dim) + def make_field( self, data_node: dace_nodes.AccessNode, @@ -578,10 +586,12 @@ def make_field( # the local dimension is converted into `ListType` data element if not isinstance(data_type.dtype, ts.ScalarType): raise ValueError(f"Invalid field type {data_type}.") - if not gtx_common.has_offset(self.offset_provider_type, local_dim.tag): + try: + self.connectivity_key_over(local_dim) + except KeyError as ex: raise ValueError( f"The provided local dimension {local_dim} does not match any offset provider type." - ) + ) from ex local_type = ts.ListType(element_type=data_type.dtype, offset_type=local_dim) field_type = ts.FieldType( dims=[dim for dim in data_type.dims if dim != local_dim], dtype=local_type @@ -839,7 +849,7 @@ def _make_array_shape_and_strides( for dim in dims: if dim.kind == gtx_common.DimensionKind.LOCAL: # for local dimension, the size is taken from the associated connectivity type - shape.append(neighbor_table_types[dim.tag].max_neighbors) + shape.append(gtx_dace_args.local_dimension_size(name, dim, neighbor_table_types)) elif gtx_dace_args.is_connectivity_identifier(name, self.offset_provider_type): # we use symbolic size for the global dimension of a connectivity shape.append(gtx_dace_args.field_size_symbol(name, dim, neighbor_table_types)) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index a28aad41c3..170fceaebd 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -1315,7 +1315,9 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if offset_type is gtx_common.ConstList: # this input argument is the result of `make_const_list` continue - offset_provider_t = self.subgraph_builder.get_offset_provider_type(offset_type.tag) + offset_provider_t = self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) input_conn_types[offset_type] = offset_provider_t @@ -1369,7 +1371,9 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if conn_type.has_skip_values: # In case the `map_list` input expressions contain skip values, we use # the connectivity-based offset provider as mask for map computation. - conn_data = gtx_dace_args.connectivity_identifier(offset_type.tag) + conn_data = gtx_dace_args.connectivity_identifier( + self.subgraph_builder.connectivity_key_over(offset_type) + ) conn_desc = self.sdfg.arrays[conn_data] conn_desc.transient = False @@ -1472,7 +1476,9 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: and input_expr.gt_dtype.offset_type is not None ) offset_type = input_expr.gt_dtype.offset_type - offset_provider_type = self.subgraph_builder.get_offset_provider_type(offset_type.tag) + offset_provider_type = self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) inp_conn = "_in" @@ -1484,7 +1490,9 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: and input_expr.gt_dtype.offset_type is not None ) offset_type = input_expr.gt_dtype.offset_type - connectivity = gtx_dace_args.connectivity_identifier(offset_type.tag) + connectivity = gtx_dace_args.connectivity_identifier( + self.subgraph_builder.connectivity_key_over(offset_type) + ) self.sdfg.arrays[connectivity].transient = False reduce_node = gtx_library_nodes.ReduceWithSkipValues( diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py index 5b0b0601fc..c614719aff 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -111,6 +111,27 @@ def field_stride_symbol( return _field_symbol(field_name, dim, "stride", offset_provider_type) +def local_dimension_size( + field_name: str, + dim: gtx_common.Dimension, + neighbor_table_types: dict[str, gtx_common.NeighborConnectivityType], +) -> int: + """ + Number of neighbors along the local dimension `dim` of the field or connectivity table. + + A connectivity table has its own neighbor count. Any other field finds it in a table over + `dim`: normally the one keyed by `dim`'s tag, but a connectivity sharing `dim` with another + one (see `NeighborConnectivity`) is keyed by its own tag, and may be the only one bound. + """ + if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is not None: + own_type = neighbor_table_types[gtx_common.from_codegen_name(m[1])] + if own_type.neighbor_dim == dim: + return own_type.max_neighbors + return neighbor_table_types[ + gtx_common.connectivity_key_over(neighbor_table_types, dim) + ].max_neighbors + + def _range_symbol_name(field_name: str, dim: gtx_common.Dimension) -> str: """Common part of the name for the range start/stop symbols.""" field_range = im.call("get_domain_range")(field_name, dim) diff --git a/tests/next_tests/definitions.py b/tests/next_tests/definitions.py index 86a4229860..eb659525ed 100644 --- a/tests/next_tests/definitions.py +++ b/tests/next_tests/definitions.py @@ -98,10 +98,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): USES_INDEX_FIELDS = "uses_index_fields" USES_LIFT = "uses_lift" USES_NEGATIVE_MODULO = "uses_negative_modulo" -USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM = "uses_offset_tag_differing_from_local_dim" -USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION = ( - "uses_offset_tag_differing_from_local_dim_in_reduction" -) USES_ORIGIN = "uses_origin" USES_REDUCE_WITH_LAMBDA = "uses_reduce_with_lambda" USES_SCAN = "uses_scan" @@ -142,10 +138,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): #: the connectivity is looked up in the offset provider by the *local dimension's* name. #: Lifted for the gtfn shift path by #1789; see #: `regression_tests/ffront_tests/test_offset_dimensions_names.py`. -OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE = ( - "'{marker}': '{backend}' looks the connectivity up by the local dimension's name," - " so it must equal the offset tag" -) # Index-only vs. consequential markers: # A `uses_*` marker only affects execution if it appears in one of the skip lists below (and thus # in `BACKEND_SKIP_TEST_MATRIX`); such a marker is "consequential" -- it applies the listed @@ -181,16 +173,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_SCAN_IN_STENCIL, XFAIL, BINDINGS_UNSUPPORTED_MESSAGE), (USES_SPARSE_FIELDS, XFAIL, UNSUPPORTED_MESSAGE), (USES_TUPLE_ITERATOR, XFAIL, UNSUPPORTED_MESSAGE), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] ) EMBEDDED_SKIP_LIST = [ @@ -202,11 +184,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): ), # we can't extract the field type from scan args (EMBEDDED_CONCAT_WHERE_INFINITE_DOMAIN, XFAIL, UNSUPPORTED_MESSAGE), (EMBEDDED_CONCAT_WHERE_NON_CONTIGUOUS_DOMAIN, XFAIL, UNSUPPORTED_MESSAGE), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] JAX_EMBEDDED_SKIP_LIST = EMBEDDED_SKIP_LIST + [ (USES_PROGRAM_WITH_SLICED_OUT_ARGUMENTS, XFAIL, UNSUPPORTED_MESSAGE), @@ -217,15 +194,7 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_TUPLES_ARGS_WITH_DIFFERENT_BUT_PROMOTABLE_DIMS, XFAIL, UNSUPPORTED_MESSAGE), (USES_CONCAT_WHERE, XFAIL, UNSUPPORTED_MESSAGE), ] -GTIR_EMBEDDED_SKIP_LIST = ROUNDTRIP_SKIP_LIST + [ - # NOTE: not in `ROUNDTRIP_SKIP_LIST`: the roundtrip backend passes this, only the - # lower-level `iterator/embedded.py` execution keys on the local dimension's name. - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), -] +GTIR_EMBEDDED_SKIP_LIST = ROUNDTRIP_SKIP_LIST GTFN_SKIP_TEST_LIST = ( COMMON_SKIP_TEST_LIST + DOMAIN_INFERENCE_SKIP_LIST @@ -236,12 +205,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_STRIDED_NEIGHBOR_OFFSET, XFAIL, BINDINGS_UNSUPPORTED_MESSAGE), # max_over broken, see https://github.com/GridTools/gt4py/issues/1289 (USES_MAX_OVER, XFAIL, UNSUPPORTED_MESSAGE), - # NOTE: only the reduction; #1789 lifted this for the shift path. - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py index 68fd18f439..9858fc0bb0 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py @@ -8,6 +8,9 @@ """A `NeighborConnectivity` declaration used directly in DSL code, on every backend.""" +import dataclasses +import typing + import numpy as np import pytest @@ -33,7 +36,7 @@ class Local(gtx.LocalDimensionIndex): ... #: A second connectivity over the same neighbor axis, bound to a different table. class V2EShared(gtx.NeighborConnectivity[V, E]): - Local = V2E.Local + Local: typing.TypeAlias = V2E.Local @pytest.fixture @@ -118,7 +121,6 @@ def testee(a: Field[Dims[E], float], out: Field[Dims[V], float]): @pytest.mark.uses_unstructured_shift -@pytest.mark.uses_offset_tag_differing_from_local_dim def test_shift_through_a_shared_local_dimension(case): @gtx.field_operator def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: @@ -128,8 +130,6 @@ def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: @pytest.mark.uses_unstructured_shift -@pytest.mark.uses_offset_tag_differing_from_local_dim -@pytest.mark.uses_offset_tag_differing_from_local_dim_in_reduction def test_reduction_through_a_shared_local_dimension(case): @gtx.field_operator def testee( @@ -143,3 +143,37 @@ def testee( testee, lambda s, a: np.sum(s * a[_table(case, V2EShared)] - a[_table(case)], axis=1), ) + + +@pytest.fixture +def case_without_owner(case): + """Only the sharing connectivity is bound: enough for a shift, which needs only its table.""" + return dataclasses.replace( + case, offset_provider={V2EShared.offset_tag: case.offset_provider[V2EShared.offset_tag]} + ) + + +@pytest.mark.uses_unstructured_shift +def test_shift_through_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2EShared[1]) + + cases.verify_with_default_data( + case_without_owner, testee, lambda a: a[_table(case_without_owner, V2EShared)[:, 1]] + ) + + +@pytest.mark.uses_unstructured_shift +def test_reduction_through_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + return neighbor_sum(s * a(V2EShared), axis=V2E.Local) + + cases.verify_with_default_data( + case_without_owner, + testee, + lambda s, a: np.sum(s * a[_table(case_without_owner, V2EShared)], axis=1), + ) diff --git a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py index 689d0d5f71..82ba0ecbd4 100644 --- a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py +++ b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py @@ -56,16 +56,8 @@ class Neigh(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) -def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases.Case: - """ - A `Case` whose offset provider holds exactly one connectivity, keyed on its tag. - - One entry per `Case` on purpose: DaCe walks every provider entry while building the - SDFG, and looks a connectivity up by its *local dimension's* name - (`gtir_to_sdfg.py`, constraint A4). A second, non-conforming entry would therefore - fail a program that does not even use it, and the cell under test would be measuring - the wrong thing. - """ +def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimension) -> cases.Case: + """A `Case` binding the same table under each of `tags`.""" mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) # NOTE: `.asnumpy()`, not `.ndarray`: under a GPU allocator the latter is a device # array, and `simple_mesh` builds the table from NumPy anyway. @@ -84,6 +76,7 @@ def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases skip_value=None, allocator=exec_alloc_descriptor.allocator, ) + for tag in tags }, default_sizes={V: mesh.num_vertices, E: mesh.num_edges}, grid_type=common.GridType.UNSTRUCTURED, @@ -93,12 +86,14 @@ def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases @pytest.fixture def case_tag_vs_variable_name(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, TaggedOffDim.tag, TaggedOffDim) + return _case(exec_alloc_descriptor, (TaggedOffDim.tag,), TaggedOffDim) @pytest.fixture def case_tag_vs_local_dim(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, "OffB", Neigh) + # NOTE: only the offset's table: a reduction over `Neigh` finds it as the table over `Neigh`, + # as for a connectivity sharing another one's local dimension. + return _case(exec_alloc_descriptor, ("OffB",), Neigh) def _neighbor_table(case: cases.Case, tag: str) -> np.ndarray: @@ -135,11 +130,9 @@ def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: # --- N3: the tag differs from the local dimension's name -------------------------- -# Lifted for the gtfn shift path by #1789; still required elsewhere, which is what -# the markers below record. +# The shape of a connectivity sharing another one's local dimension. -@pytest.mark.uses_offset_tag_differing_from_local_dim def test_shift_tag_differs_from_local_dim_name(case_tag_vs_local_dim): """ Ensure a shift works with an offset tag that differs from the local dimension's name. @@ -160,7 +153,6 @@ def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: ) -@pytest.mark.uses_offset_tag_differing_from_local_dim_in_reduction def test_reduction_tag_differs_from_local_dim_name(case_tag_vs_local_dim): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py index 1028fd987f..f104ed3b40 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py @@ -11,7 +11,7 @@ from gt4py.next import common, utils from gt4py.next.iterator import ir from gt4py.next.iterator.ir_utils import ir_makers as im -from gt4py.next.iterator.transforms.unroll_reduce import UnrollReduce, _get_partial_offset_tags +from gt4py.next.iterator.transforms.unroll_reduce import UnrollReduce, _get_partial_local_dims from gt4py.next.type_system import type_specifications as ts @@ -98,10 +98,10 @@ def reduction_if(): "reduction_if", ], ) -def test_get_partial_offsets(reduction, request): - partial_offsets = _get_partial_offset_tags(request.getfixturevalue(reduction).args) +def test_get_partial_local_dims(reduction, request): + partial_local_dims = _get_partial_local_dims(request.getfixturevalue(reduction).args) - assert set(partial_offsets) == {Dim.tag} + assert set(partial_local_dims) == {Dim} def _expected(red, max_neighbors, has_skip_values, shifted_arg=0): From 432245e3c42277f1edfc585fca60d68a5dcdc54c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 03:19:55 +0200 Subject: [PATCH 08/13] fix[next]: address review of shared local dimensions in backends - DaCe if/scan/concat_where/const-list sites find the table over the local dimension - connectivity_key_over takes a tag, avoids resolve, and picks deterministically - embedded map_list compares lists by local dimension, not by offset - a sharer must have its owner's origin; counts are compared one by one - ADR 0029: sharers must have the owner's neighbor structure --- .../ADRs/next/0029-Connectivities_As_Types.md | 23 +++--- src/gt4py/next/common.py | 71 ++++++++++++------- src/gt4py/next/iterator/embedded.py | 21 +++--- .../lowering/gtir_to_sdfg_concat_where.py | 4 +- .../dace/lowering/gtir_to_sdfg_lambda.py | 6 +- .../dace/lowering/gtir_to_sdfg_primitives.py | 4 +- .../dace/lowering/gtir_to_sdfg_scan.py | 4 +- .../test_neighbor_connectivity.py | 20 ++++++ .../test_with_toy_connectivity.py | 50 ++++++++++++- .../dace_tests/test_dace_utils.py | 31 ++++++++ .../unit_tests/test_neighbor_connectivity.py | 58 ++++++++++++++- 11 files changed, 240 insertions(+), 52 deletions(-) diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index 150976307b..e2b433fcd6 100644 --- a/docs/development/ADRs/next/0029-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -124,15 +124,20 @@ whose tag is the connectivity's `offset_tag`: existing backends need no change. - **its own tag**, `C2CE.tag`, for a connectivity that shares another one's local dimension, since the local dimension's tag already names the owner's table. - Shifts find the table by that tag; reductions and sparse arguments still find - the owner's table by the local dimension's tag, for its neighbor structure, so - the owner has to be bound too. - `V2E.Local` inside DSL code types as that local dimension, and - `FieldOffset.Local` names the same thing on a legacy offset, so the spelling - works for both. The other frontend touch points treat the class like the - `FieldOffset` it derives: grid-type deduction (`transform_utils`, `past_to_itir`) - counts it as unstructured, and embedded `premap` accepts it. `V2E[i]` subscripts the - metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. + Shifts find the table by that tag. Reductions and sparse arguments know only + the local dimension, and take its neighbor count and skip values from a table + over it (`common.connectivity_key_over`): the owner's if bound, else the + sharer with the smallest tag. Connectivities sharing a local dimension must + therefore have the same neighbor *structure* — the same count, and a skip value + at the same positions — which is what sharing a neighbor axis means; + `check_offset_provider` enforces it for the tables it is given. + +`V2E.Local` inside DSL code types as that local dimension, and +`FieldOffset.Local` names the same thing on a legacy offset, so the spelling +works for both. The other frontend touch points treat the class like the +`FieldOffset` it derives: grid-type deduction (`transform_utils`, `past_to_itir`) +counts it as unstructured, and embedded `premap` accepts it. `V2E[i]` subscripts the +metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. ## Consequences diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 58729c4184..dd36f16e0d 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1544,35 +1544,45 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid def connectivity_key_over( - offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension + offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension | Tag ) -> str: """ - The key of a bound connectivity whose local dimension is `local_dim`. + The key of a bound connectivity whose local dimension is `local_dim` (a dimension or its tag). Neighbor reductions and sparse fields know only their local dimension, and use its table for the neighbor count and the skip values. That is the table keyed by the local dimension's - tag, i.e. its owner's, if bound. Otherwise it is a connectivity *sharing* the local dimension - (see `NeighborConnectivity`), keyed by its own tag, which has the same neighbor structure. + tag, i.e. its owner's, if bound. Otherwise it is one of the connectivities *sharing* the local + dimension (see `NeighborConnectivity`), each keyed by its own tag; the smallest key is taken, + so the choice does not depend on the order of the provider. Connectivities sharing a local + dimension have the same neighbor structure (see `check_offset_provider`), so which one does + not matter. Raises: KeyError: If no bound connectivity has `local_dim` as its local dimension. """ - if local_dim.tag in offset_provider: - return local_dim.tag - for key, connectivity in offset_provider.items(): - if isinstance(connectivity, NeighborConnectivityType): - neighbor_dim = connectivity.neighbor_dim - elif is_neighbor_table(connectivity): - neighbor_dim = connectivity.domain.dims[1] - else: - continue - if neighbor_dim is local_dim: - assert isinstance(key, str) - return key - raise KeyError( - f"No connectivity over the local dimension '{local_dim.tag}' is bound in the offset" - f" provider, which has {sorted(map(str, offset_provider))}." - ) + local_tag = local_dim if isinstance(local_dim, str) else local_dim.tag + if local_tag in offset_provider: + return local_tag + candidates = [ + key + for key, connectivity in offset_provider.items() + if (neighbor_dim := _neighbor_dim_of(connectivity)) is not None + and neighbor_dim.tag == local_tag + ] + if not candidates: + raise KeyError( + f"No connectivity over the local dimension '{local_tag}' is bound in the offset" + f" provider, which has {sorted(map(str, offset_provider))}." + ) + return min(candidates) + + +def _neighbor_dim_of(connectivity: Any) -> Optional[Dimension]: + if isinstance(connectivity, NeighborConnectivityType): + return connectivity.neighbor_dim + if is_neighbor_table(connectivity): + return connectivity.domain.dims[1] + return None def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: str) -> bool: @@ -2230,14 +2240,23 @@ def __init_subclass__( if local.owner is not None and local.owner.tag != cls.tag: # Sharing another connectivity's local dimension: the neighbor structure is the # owner's, including its counts, and the sharing connectivity is named by its own tag. - if (max_neighbors, min_neighbors) != (None, None) and ( - max_neighbors, - min_neighbors, - ) != (local.max_neighbors, local.min_neighbors): + owner_name = local.owner.__qualname__ + if origin is not local.owner.origin: raise TypeError( - f"'{name}' shares the local dimension of '{local.owner.__qualname__}', whose" - " neighbor counts are declared by its owner." + f"'{name}' cannot share the local dimension of '{owner_name}': it has origin" + f" '{origin}', but the neighbors of '{owner_name}' are those of" + f" '{local.owner.origin}'." ) + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("min_neighbors", min_neighbors), + ): + if count is not None and count != getattr(local, count_name): + raise TypeError( + f"'{name}': '{count_name}={count}' contradicts the local dimension it" + f" shares with '{owner_name}', which declares" + f" {count_name}={getattr(local, count_name)}." + ) cls.origin, cls.codomain = origin, codomain return for count_name, count in ( diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 985bc139cc..3adb6922aa 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -573,7 +573,7 @@ def execute_shift( # keyed by a connectivity sharing it (see `common.connectivity_key_over`). offset_implementation = common.get_offset( offset_provider, - common.connectivity_key_over(offset_provider, common.resolve(tag)), + common.connectivity_key_over(offset_provider, tag), ) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().source_dim @@ -1003,9 +1003,10 @@ def field_getitem(self, named_indices: NamedFieldIndices) -> Any: def field_setitem(self, named_indices: NamedFieldIndices, value: Any): if isinstance(self._ndarrayfield, common.MutableField): if isinstance(value, _List): + local_tag = value.local_dim.tag for i, v in enumerate(value): # type:ignore[var-annotated, arg-type] self._ndarrayfield[ - self._translate_named_indices({**named_indices, value.local_dim.tag: i}) + self._translate_named_indices({**named_indices, local_tag: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ @@ -1480,12 +1481,14 @@ def list_get(i, lst: _List[Optional[DT]]) -> Optional[DT] | Undefined: def _get_offset(*lists: _List | _ConstList) -> Optional[runtime.Offset]: - offsets = set((lst.offset for lst in lists if hasattr(lst, "offset"))) - if len(offsets) == 0: + neighbor_lists = [lst for lst in lists if isinstance(lst, _List)] + if len(neighbor_lists) == 0: return None - if len(offsets) == 1: - return offsets.pop() - raise AssertionError("All lists must have the same offset.") + # NOTE: compared by local dimension, not by offset: a connectivity and one sharing its local + # dimension build lists along the same axis. + if len({lst.local_dim for lst in neighbor_lists}) != 1: + raise AssertionError("All lists must run along the same local dimension.") + return neighbor_lists[0].offset @builtins.map_list.register(EMBEDDED) @@ -1540,9 +1543,7 @@ def deref(self) -> Any: assert offset_provider is not None # NOTE: `list_offset` is the local dimension's tag; the table over it may be keyed by a # connectivity sharing it (see `common.connectivity_key_over`). - connectivity_key = common.connectivity_key_over( - offset_provider, common.resolve(self.list_offset) - ) + connectivity_key = common.connectivity_key_over(offset_provider, self.list_offset) connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py index 737648b911..60827eeedc 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py @@ -254,7 +254,9 @@ def translate_concat_where( local_dim = node.type.dtype.offset_type assert local_dim is not None dtype = gtx_dace_args.as_dace_type(node.type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type(local_dim.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(local_dim) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) local_idx = gtx_common.order_dimensions([*output_dims, local_dim]).index(local_dim) output_shape.insert(local_idx, offset_provider_type.max_neighbors) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 170fceaebd..9ba625ea26 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -762,7 +762,9 @@ def _visit_if_branch_arg( local_dim = arg.gt_dtype.offset_type assert local_dim is not None assert isinstance( - self.subgraph_builder.get_offset_provider_type(local_dim.tag), + self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(local_dim) + ), gtx_common.NeighborConnectivityType, ) # find position of the local dimension in the field layout @@ -1436,7 +1438,7 @@ def _broadcast_const_list( ) -> ValueExpr: assert list_type.offset_type is not None offset_provider_t = self.subgraph_builder.get_offset_provider_type( - list_type.offset_type.tag + self.subgraph_builder.connectivity_key_over(list_type.offset_type) ) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) local_size = offset_provider_t.max_neighbors diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index e265377a45..50a445146c 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -325,7 +325,9 @@ def _construct_if_branch_output( assert out_type.dtype.offset_type is not None assert isinstance(out_type.dtype.element_type, ts.ScalarType) dtype = gtx_dace_args.as_dace_type(out_type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type(out_type.dtype.offset_type.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(out_type.dtype.offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) shape = [*shape, offset_provider_type.max_neighbors] diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index 18b48987cb..b158097a51 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -384,7 +384,9 @@ def get_scan_output_shape( assert isinstance(scan_init_data.gt_type, ts.ListType) assert scan_init_data.gt_type.offset_type offset_type = scan_init_data.gt_type.offset_type - offset_provider_type = sdfg_builder.get_offset_provider_type(offset_type.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) list_size = offset_provider_type.max_neighbors return [scan_column_size, dace.symbolic.SymExpr(list_size)] diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py index 9858fc0bb0..bce54f4e6a 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py @@ -177,3 +177,23 @@ def testee( testee, lambda s, a: np.sum(s * a[_table(case_without_owner, V2EShared)], axis=1), ) + + +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_if_stmts +def test_if_over_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee(a: Field[Dims[E], float], flag: bool) -> Field[Dims[V], float]: + if flag: + s = a(V2EShared) + else: + s = a(V2EShared) * 2.0 + return neighbor_sum(s, axis=V2E.Local) + + cases.verify_with_default_data( + case_without_owner, + testee, + lambda a, flag: ( + np.sum(a[_table(case_without_owner, V2EShared)], axis=1) * (1.0 if flag else 2.0) + ), + ) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py index fdf8cf5114..2490d5ab91 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py @@ -23,7 +23,7 @@ reduce, shift, ) -from gt4py.next.iterator.runtime import fundef +from gt4py.next.iterator.runtime import fundef, offset from gt4py.next.program_processors.runners import gtfn from next_tests.toy_connectivity import ( @@ -434,3 +434,51 @@ def test_sparse_shifted_stencil_reduce(program_processor): if validate: assert np.allclose(out.asnumpy(), ref) + + +#: A second connectivity over `V2E`'s local dimension, under its own name; bound to the table +#: with its columns reversed. +V2E_SHARED = offset("V2EShared") + + +v2e_shared_arr = np.ascontiguousarray(v2e_arr[:, ::-1]) +v2e_shared_conn = gtx.as_connectivity( + domain={Vertex: v2e_shared_arr.shape[0], V2EDim: v2e_shared_arr.shape[1]}, + codomain=Edge, + data=v2e_shared_arr, +) + + +@fundef +def shift_through_sharer(in_edges): + return deref(shift(V2E_SHARED, 1)(in_edges)) + + +@fundef +def owner_times_sharer(in_edges): + return reduce(plus, 0)( + map_list(multiplies)(neighbors(V2E_SHARED, in_edges), neighbors(V2E, in_edges)) + ) + + +@pytest.mark.parametrize( + "stencil, ref", + [ + (shift_through_sharer, v2e_shared_arr[:, 1]), + (owner_times_sharer, np.sum(v2e_shared_arr * v2e_arr, axis=1)), + ], +) +def test_connectivity_sharing_a_local_dimension(program_processor, stencil, ref): + program_processor, validate = program_processor + inp = edge_index_field() + out = gtx.as_field([Vertex], np.zeros([9], dtype=inp.dtype)) + + run_processor( + stencil[{Vertex: range(0, 9)}], + program_processor, + inp, + out=out, + offset_provider={V2EDim.tag: v2e_conn, V2E_SHARED.value: v2e_shared_conn}, + ) + if validate: + assert np.allclose(out.asnumpy(), ref) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py index 6fdec5a21f..50141db06b 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py @@ -18,3 +18,34 @@ def test_safe_replace_symbolic(): assert gtir_to_sdfg_utils.safe_replace_symbolic( dace.symbolic.pystr_to_symbolic("x*x + y"), symbol_mapping={"x": "y", "y": "x"} ) == dace.symbolic.pystr_to_symbolic("y*y + x") + + +def test_local_dimension_size(): + import numpy as np + + from gt4py._core import definitions as core_defs + from gt4py.next import common + from gt4py.next.program_processors.runners.dace import sdfg_args + + from next_tests.toy_connectivity import V2EDim, Vertex, Edge + + def conn_type(max_neighbors: int) -> common.NeighborConnectivityType: + return common.NeighborConnectivityType( + domain=(Vertex, V2EDim), + codomain=Edge, + skip_value=None, + dtype=core_defs.dtype(np.int32), + max_neighbors=max_neighbors, + ) + + sharer_tag = "some.module.V2EShared" + table_types = {V2EDim.tag: conn_type(4), sharer_tag: conn_type(4)} + # a field finds the size in the table keyed by the local dimension + assert sdfg_args.local_dimension_size("a_field", V2EDim, table_types) == 4 + # a connectivity array has its own + conn_array = sdfg_args.connectivity_identifier(sharer_tag) + assert sdfg_args.local_dimension_size(conn_array, V2EDim, {sharer_tag: conn_type(4)}) == 4 + # a field over a local dimension bound only through a sharing connectivity + assert sdfg_args.local_dimension_size("a_field", V2EDim, {sharer_tag: conn_type(4)}) == 4 + with pytest.raises(KeyError, match="No connectivity over the local dimension"): + sdfg_args.local_dimension_size("a_field", V2EDim, {}) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index c5f3a78743..3ff9bcc055 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -146,7 +146,14 @@ class Local(DimensionIndex, kind=DimensionKind.LOCAL): ... class C(NeighborConnectivity[Vertex, Edge], max_neighbors=5): Local: typing.TypeAlias = V2E.Local """, - "counts are declared by its owner", + "contradicts the local dimension it shares with 'V2E'", + ), + ( + """ + class C(NeighborConnectivity[Edge, Edge]): + Local = V2E.Local + """, + "cannot share the local dimension of 'V2E'", ), ( """ @@ -243,6 +250,15 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): assert shared.offset_tag == shared.tag assert shared.__gt_type__().tag == shared.tag + def test_sharing_with_consistent_counts(self): + ns = _declare( + """ + class V2EShared4(NeighborConnectivity[Vertex, Edge], max_neighbors=4): + Local = V2E.Local + """ + ) + assert ns["V2EShared4"].Local is V2E.Local + def test_non_integer_index(self): with pytest.raises(TypeError, match="indexed by an integer"): V2E[Vertex] @@ -475,6 +491,46 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): common.local_dimension_of(NeighborConnectivity) +class TestConnectivityKeyOver: + def _type(self, connectivity): + return _table_type(domain=(connectivity.origin, common.local_dimension_of(connectivity))) + + def test_owner_is_preferred(self): + ns = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + ) + shared = ns["V2EShared"] + provider = {shared.offset_tag: self._type(shared), V2E.offset_tag: self._type(V2E)} + assert common.connectivity_key_over(provider, V2E.Local) == V2E.offset_tag + assert common.connectivity_key_over(provider, V2E.Local.tag) == V2E.offset_tag + + def test_sharers_are_picked_independently_of_order(self): + ns = _declare( + """ + class SharedA(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + + class SharedB(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + ) + a, b = ns["SharedA"], ns["SharedB"] + forward = {a.offset_tag: self._type(a), b.offset_tag: self._type(b)} + backward = dict(reversed(forward.items())) + assert ( + common.connectivity_key_over(forward, V2E.Local) + == common.connectivity_key_over(backward, V2E.Local) + == min(a.offset_tag, b.offset_tag) + ) + + def test_nothing_bound(self): + with pytest.raises(KeyError, match="No connectivity over the local dimension"): + common.connectivity_key_over({E2V.offset_tag: self._type(E2V)}, V2E.Local) + + def test_the_const_list_dimension_cannot_be_adopted(): with pytest.raises(TypeError, match="cannot adopt"): _declare( From b0c303983b010a3a36e159927a3d0e27f435a0a4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 24 Sep 2026 00:21:53 +0200 Subject: [PATCH 09/13] refactor[next]: NeighborConnectivity[Domain, Codomain], NeighborTableType and ts.ShiftType - 'NeighborConnectivity[Origin, Codomain]' becomes '[Domain, Codomain]', and the class attribute 'origin' becomes 'domain': a bound table is a field over (Domain, Local) with values in Codomain. 'origin' already names a buffer origin ('__gt_origin__'). - 'NeighborConnectivityType' becomes 'NeighborTableType': the declaration the table is bound to, dtype, skip value and max_neighbors, with domain and codomain derived from the declaration. It is built from the offset-provider key ('offset_provider_to_type', 'check_neighbor_table', which now returns it), since a sharer's table looks like its owner's; local dimensions record their sharers for that. A table under a key no declaration answers to keeps its structural 'ConnectivityType', which is also what 'NeighborTable.__gt_type__()' returns now. 'OffsetProviderType' becomes 'TableTypes'. - 'ts.OffsetType(source, target)' becomes 'ts.ShiftType(codomain, domain)', printed 'Shift[: -> ]'. - ADR 0029 describes the result; ADR 0019 names the new record. --- .../ADRs/next/0019-Connectivities.md | 8 +- .../ADRs/next/0029-Connectivities_As_Types.md | 80 ++++++-- src/gt4py/next/common.py | 193 ++++++++++++------ src/gt4py/next/ffront/decorator.py | 15 +- src/gt4py/next/ffront/fbuiltins.py | 19 +- .../ffront/foast_passes/type_deduction.py | 52 ++--- src/gt4py/next/ffront/foast_to_gtir.py | 8 +- src/gt4py/next/fingerprinting.py | 2 +- src/gt4py/next/iterator/embedded.py | 25 +-- .../next/iterator/ir_utils/domain_utils.py | 6 +- src/gt4py/next/iterator/runtime.py | 6 +- .../iterator/transforms/collapse_tuple.py | 2 +- .../concat_where/expand_tuple_args.py | 2 +- src/gt4py/next/iterator/transforms/cse.py | 2 +- .../transforms/dead_code_elimination.py | 2 +- .../iterator/transforms/expand_tuple_maps.py | 2 +- .../iterator/transforms/fuse_as_fieldop.py | 6 +- .../next/iterator/transforms/global_tmps.py | 2 +- .../next/iterator/transforms/infer_domain.py | 10 +- .../transforms/inline_dynamic_shifts.py | 4 +- .../next/iterator/transforms/inline_scalar.py | 2 +- .../next/iterator/transforms/pass_manager.py | 6 +- .../transforms/prune_empty_concat_where.py | 6 +- .../next/iterator/transforms/unroll_reduce.py | 12 +- .../next/iterator/type_system/inference.py | 8 +- .../iterator/type_system/type_synthesizer.py | 55 +++-- src/gt4py/next/otf/arguments.py | 4 +- src/gt4py/next/otf/compiled_program.py | 12 +- src/gt4py/next/otf/options.py | 2 +- .../codegens/gtfn/gtfn_module.py | 14 +- .../codegens/gtfn/itir_to_gtfn_ir.py | 20 +- .../runners/dace/lowering/gtir_to_sdfg.py | 10 +- .../lowering/gtir_to_sdfg_concat_where.py | 2 +- .../dace/lowering/gtir_to_sdfg_lambda.py | 36 ++-- .../dace/lowering/gtir_to_sdfg_primitives.py | 2 +- .../dace/lowering/gtir_to_sdfg_scan.py | 2 +- .../runners/dace/sdfg_args.py | 26 +-- .../runners/dace/workflow/translation.py | 8 +- src/gt4py/next/type_system/type_info.py | 30 +-- .../next/type_system/type_specifications.py | 22 +- .../next/type_system/type_translation.py | 2 +- .../integration_tests/cases_utils.py | 4 +- .../instrumentation_tests/test_hooks.py | 2 +- .../test_strided_offset_provider.py | 6 +- .../test_offset_dimensions_names.py | 2 +- .../ir_utils_tests/test_domain_utils.py | 2 +- .../iterator_tests/test_runtime_domain.py | 11 +- .../transforms_tests/test_unroll_reduce.py | 13 +- .../dace_tests/test_dace_utils.py | 12 +- .../unit_tests/test_neighbor_connectivity.py | 124 +++++++++-- 50 files changed, 557 insertions(+), 346 deletions(-) diff --git a/docs/development/ADRs/next/0019-Connectivities.md b/docs/development/ADRs/next/0019-Connectivities.md index 21827941bf..f22a13e039 100644 --- a/docs/development/ADRs/next/0019-Connectivities.md +++ b/docs/development/ADRs/next/0019-Connectivities.md @@ -31,7 +31,7 @@ We update and introduce the following concepts **NeighborTable** is a _GatherConnectivity_ that is a 2D mapping of the N neighbors of a Location A to a Location B, backed by a buffer. -**ConnectivityType**, **NeighborConnectivityType** contains all information that is needed for compilation. +**ConnectivityType**, **NeighborTableType** contain all information that is needed for compilation. A `NeighborTableType` is the type of a table bound to a `NeighborConnectivity` declaration (ADR 0029). ### Full definitions @@ -48,7 +48,7 @@ Embedded execution of iterator (local) view supports only `NeighborTable`s. ### IR transformations and compiled backends -All transformations and code-generation should use `ConnectivityType`, not the `Connectivity` which contains the runtime mapping. +All transformations and code-generation should use `ConnectivityType` / `NeighborTableType`, not the `Connectivity` which contains the runtime mapping. Note, currently the `global_tmps` pass uses runtime information, therefore this is not strictly enforced. @@ -60,3 +60,7 @@ The only supported `Connectivity`s in compiled backends (currently) are `Neighbo - Removed the abstract `NeighborConnectivity` concept; `NeighborTable` is now the single neighbor-connectivity concept (there is no non-buffer-backed neighbor connectivity in use). - Added `GatherConnectivity` (a `Connectivity` whose `premap` rearranges data via a gather), which the embedded field-view `premap` dispatches on. It replaces the former `ConnectivityKind` flag and unifies the previous reshuffling/remapping `premap` implementations into a single advanced-index gather. + +### 2026-09-24 + +- `NeighborConnectivityType` is renamed `NeighborTableType` and typed by the `NeighborConnectivity` declaration its table is bound to; a `NeighborTable`'s own `__gt_type__()` is the structural `ConnectivityType` (ADR 0029). diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index e2b433fcd6..a08b1555f7 100644 --- a/docs/development/ADRs/next/0029-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -7,7 +7,7 @@ tags: [] - **Status**: proposed - **Authors**: Enrique González Paredes (@egparedes) - **Created**: 2026-09-21 -- **Updated**: 2026-09-21 +- **Updated**: 2026-09-24 A neighbor connectivity is declared as a **class**, and its local dimension as a class **nested** in it: @@ -43,9 +43,10 @@ and skip values were never checked against the `FieldOffset` declaration. ### The declaration -- `NeighborConnectivity[Origin, Codomain]` is a PEP 695 generic whose subclasses - are declarations: for each `Origin` element, a list of `Codomain` neighbors. Its - metaclass, `ConnectivityMeta`, forbids instantiation. +- `NeighborConnectivity[Domain, Codomain]` is a PEP 695 generic whose subclasses + are declarations: for each `Domain` element, a list of `Codomain` neighbors. Its + metaclass, `ConnectivityMeta`, forbids instantiation. The two dimensions are + the class attributes `V2E.domain` and `V2E.codomain`. - The local dimension is the nested class `Local`, a subclass of `LocalDimensionIndex`. Declaring it is required, and `NeighborConnectivity` sets `Local.owner` to the connectivity when the class is created. A local @@ -60,21 +61,59 @@ and skip values were never checked against the `FieldOffset` declaration. `Local: TypeAlias = C2E.Local`. This is the flattened sparse pattern, e.g. cell-to-cell-edge (`C2CE: Cell -> CellEdge`) indexing the same neighbor axis as `C2E`, so that its results combine with `C2E`-shaped sparse fields. The owner stays `C2E`, and the - neighbor counts and skip-value structure are the owner's. + neighbor counts and skip-value structure are the owner's. A sharer must have + the owner's domain; its codomain is free. The local dimension records its + sharers (`Local.sharers`) as it records its owner. - `max_neighbors` and `min_neighbors` are optional class keywords, not type parameters: Python has no integer type parameters, and nothing static needs the count. A declared count is a constraint on the bound table; an undeclared one is taken from the table. `min_neighbors < max_neighbors` means that the table must use skip values. - `common.check_neighbor_table(V2E, table)` checks a table, or just its type - (which is all an ahead-of-time compilation has), against the declaration: - the domain is `(Origin, V2E.Local)`, the codomain is `Codomain`, the dtype is - integral, and the neighbor counts and skip values agree. Skip values are - checked on the table's type: a table with a `skip_value` counts as having skip - values whether or not an entry uses it. The check is explicit for as long as - offset providers are keyed by tag strings, since nothing then connects a - provider entry to a declaration; it becomes automatic with class-keyed - providers. + (which is all an ahead-of-time compilation has), against the declaration, and + returns the table's `NeighborTableType` (below): the domain is + `(Domain, V2E.Local)`, the codomain is `Codomain`, the dtype is integral, and + the neighbor counts and skip values agree. Skip values are checked on the + table's type: a table with a `skip_value` counts as having skip values whether + or not an entry uses it. + +`Domain` and `Codomain` name the two index spaces the declaration maps between. +A bound table is a field over `(Domain, Local)` with values in `Codomain`: the +table's domain is the declaration's domain extended by the local axis, which is +the same use of the word as `Connectivity.domain` and +`CartesianConnectivity.domain_dim`. "Origin" would have been the other natural +name for the first dimension, but gt4py already uses it for the start of a +buffer (`__gt_origin__`). + +### The type of a bound table + +Transformations and code generation see types, never tables (ADR 0019). The type +of a table bound to a declaration is a `common.NeighborTableType`: +`connectivity` (the declaration), `dtype`, `skip_value` and `max_neighbors`. Its +`domain` and `codomain` are derived from the declaration, +`(connectivity.domain, local_dimension_of(connectivity))` and +`connectivity.codomain`, so they cannot disagree with it. The mapping from +offset-provider keys to these records is `common.TableTypes`, and it can be +given instead of the tables for ahead-of-time compilation. + +A table cannot tell which declaration it is bound to: the table of a sharer +(`C2CE`) has the same domain as its owner's (`C2E`), with another codomain. So a +`NeighborTableType` is built where a table is bound, from its offset-provider key: +`check_neighbor_table(C2CE, table)`, or `offset_provider_to_type`, which finds +the declaration whose `offset_tag` is the key among the owner and the sharers of +the table's local dimension. `NeighborTable.__gt_type__()` returns only what the +table knows, the structural `common.ConnectivityType` (domain, codomain, dtype, +skip value). + +A table bound under a key that no declaration answers to -- hand-written IR +names its offsets by plain strings -- has no declaration. Its +`NeighborTableType` then has the table's structural `ConnectivityType` as its +`connectivity`, and `domain` and `codomain` are read from that. This keeps the +IR level, which does not know declarations, working unchanged. + +A `NeighborTableType` is fingerprinted through its fields, so the declaration +takes part in the fingerprint of everything compiled for it: the owner's and a +sharer's tables, identical as tables, produce different artifact keys. ### `NeighborConnectivity` is not a `Connectivity` @@ -91,7 +130,7 @@ A separate root would force every `type[DimensionIndex]` annotation in the tree would then accept local dimensions wherever a primary one is meant anyway. The tree already distinguishes local dimensions by a runtime `kind` check, so it keeps doing so; generic constructors whose parameter must be a primary dimension -(`NeighborConnectivity[Origin, Codomain]`, `Staggered[D]`) check it at runtime. +(`NeighborConnectivity[Domain, Codomain]`, `Staggered[D]`) check it at runtime. ### `Local` is not annotated anywhere @@ -115,8 +154,10 @@ treating it as a type. ### Frontend integration A declaration is typed like the `FieldOffset` it replaces: `V2E.__gt_type__()` -is the `ts.OffsetType` of the derived offset `(Codomain -> (Origin, Local))`, -whose tag is the connectivity's `offset_tag`: +is a `ts.ShiftType`, which takes a field over the codomain to one over the +domain, `Shift[: Edge -> (Vertex, V2E.Local)]`. `V2E[i]` has the domain +`(Vertex,)`, and so does a Cartesian shift `KDim + 1`, over `KDim` and without a +tag. The tag is the connectivity's `offset_tag`: - **the local dimension's tag**, `V2E.Local.tag`, for the connectivity that declares it. This is the single string that shifts, neighbor reductions and @@ -129,8 +170,7 @@ whose tag is the connectivity's `offset_tag`: over it (`common.connectivity_key_over`): the owner's if bound, else the sharer with the smallest tag. Connectivities sharing a local dimension must therefore have the same neighbor *structure* — the same count, and a skip value - at the same positions — which is what sharing a neighbor axis means; - `check_offset_provider` enforces it for the tables it is given. + at the same positions — which is what sharing a neighbor axis means. `V2E.Local` inside DSL code types as that local dimension, and `FieldOffset.Local` names the same thing on a legacy offset, so the spelling @@ -144,8 +184,10 @@ metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, - An unstructured connectivity is spelled once. The provider key, the offset tag and the local dimension are all derived from the declaration. - A table bound to a connectivity can be checked against its declaration. -- `V2E.Local` in DSL code is resolved from the offset type, because the type of +- `V2E.Local` in DSL code is resolved from the shift type, because the type of `V2E` is not the class. +- Code generation sees which declaration a table is bound to, not only its + shape; a table without a declaration is typed by its structure. - A declaration is fingerprinted by its name *and* its declared dimensions and counts, so redefining it under the same name (e.g. re-running a notebook cell) does not reuse artifacts compiled for the old declaration. diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index dd36f16e0d..35ae15d33d 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1301,20 +1301,44 @@ def has_skip_values(self) -> bool: @dataclasses.dataclass(frozen=True) -class NeighborConnectivityType(ConnectivityType): - # NOTE: partly encoded in the local dimension since ADR 0029: a `LocalDimensionIndex` carries - # `max_neighbors` / `min_neighbors` where the declaration states them, and this record is - # checked against them (`check_neighbor_table`). It stays the *bound* count, which a - # declaration may leave to the table. +class NeighborTableType: + """ + The type of a neighbor table bound to a connectivity: what transformations and code generation + see instead of the table (ADR 0019). + + `connectivity` is the `NeighborConnectivity` declaration the table is bound to. It determines + the table's `domain` -- the declaration's domain extended by its local dimension -- and its + `codomain`. A table alone cannot name its declaration: one sharing another connectivity's + local dimension has a table over the same domain as the owner's, with another codomain. So the + record is built where a table is bound, from its offset-provider key (`offset_provider_to_type`, + `check_neighbor_table`), or given directly for ahead-of-time compilation. + + A table bound under a name that no declaration answers to, as hand-written IR binds them, has + no declaration: `connectivity` is then the table's own structural `ConnectivityType`, which is + also what `NeighborTable.__gt_type__()` returns. + """ + + connectivity: type[NeighborConnectivity] | ConnectivityType + dtype: core_defs.DType + skip_value: Optional[core_defs.IntegralScalar] + #: The table's number of entries per element. A declaration may leave it to the table; where + #: it states one, `check_neighbor_table` checks the table against it. max_neighbors: int @property - def source_dim(self) -> Dimension: - return self.domain[0] + def domain(self) -> tuple[Dimension, Dimension]: + if isinstance(self.connectivity, ConnectivityType): + first, second = self.connectivity.domain + return (first, second) + return (self.connectivity.domain, local_dimension_of(self.connectivity)) @property - def neighbor_dim(self) -> Dimension: - return self.domain[1] + def codomain(self) -> Dimension: + return self.connectivity.codomain + + @property + def has_skip_values(self) -> bool: + return self.skip_value is not None @runtime_checkable @@ -1333,21 +1357,14 @@ def codomain(self) -> DimT_co: """ def __gt_type__(self) -> ConnectivityType: - if is_neighbor_table(self): - return NeighborConnectivityType( - domain=self.domain.dims, - codomain=self.codomain, - dtype=self.dtype, - skip_value=self.skip_value, - max_neighbors=self.ndarray.shape[1], - ) - else: - return ConnectivityType( - domain=self.domain.dims, - codomain=self.codomain, - dtype=self.dtype, - skip_value=self.skip_value, - ) + # NOTE: structural, also for a neighbor table: the table cannot tell which declaration it + # is bound to, so its `NeighborTableType` is built from its offset-provider key. + return ConnectivityType( + domain=self.domain.dims, + codomain=self.codomain, + dtype=self.dtype, + skip_value=self.skip_value, + ) @abc.abstractmethod def inverse_image(self, image_range: UnitRange | NamedRange) -> Sequence[NamedRange]: ... @@ -1478,8 +1495,7 @@ def _connectivity( @runtime_checkable class NeighborTable(Connectivity, Protocol): - # TODO(havogt): work towards encoding this properly in the type - def __gt_type__(self) -> NeighborConnectivityType: ... + def __gt_type__(self) -> ConnectivityType: ... @property def ndarray(self) -> core_defs.NDArrayObject: @@ -1501,11 +1517,12 @@ def is_neighbor_table(obj: Any) -> TypeGuard[NeighborTable]: OffsetProviderElem: TypeAlias = NeighborTable -OffsetProviderTypeElem: TypeAlias = NeighborConnectivityType -# Note: `OffsetProvider` and `OffsetProviderType` should not be accessed directly, +# Note: `OffsetProvider` and `TableTypes` should not be accessed directly, # use the `get_offset` and `get_offset_type` functions instead. OffsetProvider: TypeAlias = Mapping[Tag, OffsetProviderElem] -OffsetProviderType: TypeAlias = Mapping[Tag, OffsetProviderTypeElem] +#: The types of an offset provider's tables, under the same keys: what transformations and code +#: generation see instead of the tables (ADR 0019). +TableTypes: TypeAlias = Mapping[Tag, NeighborTableType] def is_offset_provider(obj: Any) -> TypeGuard[OffsetProvider]: @@ -1514,25 +1531,54 @@ def is_offset_provider(obj: Any) -> TypeGuard[OffsetProvider]: return all(isinstance(el, OffsetProviderElem) for el in obj.values()) -def is_offset_provider_type(obj: Any) -> TypeGuard[OffsetProviderType]: +def is_table_types(obj: Any) -> TypeGuard[TableTypes]: if not isinstance(obj, Mapping): return False - return all(isinstance(el, OffsetProviderTypeElem) for el in obj.values()) + return all(isinstance(el, NeighborTableType) for el in obj.values()) -def offset_provider_to_type( - offset_provider: OffsetProvider | OffsetProviderType, -) -> OffsetProviderType: +def offset_provider_to_type(offset_provider: OffsetProvider | TableTypes) -> TableTypes: + """The types of an offset provider's tables, each typed by the declaration its key names.""" return { - k: v.__gt_type__() if isinstance(v, Connectivity) else v for k, v in offset_provider.items() + key: value if isinstance(value, NeighborTableType) else _neighbor_table_type(key, value) + for key, value in offset_provider.items() } +def _unbound_table_type(table: NeighborTable) -> NeighborTableType: + structure = table.__gt_type__() + return NeighborTableType( + connectivity=structure, + dtype=structure.dtype, + skip_value=structure.skip_value, + max_neighbors=len(table.domain[1].unit_range), + ) + + +def _neighbor_table_type(key: Tag, table: NeighborTable) -> NeighborTableType: + """ + The type of `table` bound under `key`: typed by the declaration `key` is the `offset_tag` of. + + The declaration is found through the table's local dimension, which knows its owner and the + connectivities sharing it; a key none of them answers to leaves the table undeclared. + """ + table_type = _unbound_table_type(table) + local = table_type.domain[1] + if not (isinstance(local, DimensionMeta) and issubclass(local, LocalDimensionIndex)): + return table_type + # NOTE: the most recent sharer first: a redefined declaration (a re-run notebook cell) is + # appended again under the same tag. + for connectivity in (local.owner, *reversed(local.sharers)): + if connectivity is not None and connectivity.offset_tag == key: + return check_neighbor_table(connectivity, table_type) + return table_type + + def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProviderElem: """ - Get the `OffsetProviderElem` or `OffsetProviderTypeElem` for the given `offset` string. + Get the `OffsetProviderElem` or `NeighborTableType` for the given `offset` string. - Note: All accesses of `OffsetProvider` or `OffsetProviderType` should go through this function. + Note: All accesses of `OffsetProvider` or `TableTypes` should go through this function. """ # TODO(havogt): Once we have a custom class for `OffsetProvider`, we can absorb this functionality into it. if offset_tag not in offset_provider: @@ -1540,11 +1586,11 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid return offset_provider[offset_tag] # TODO return a valid dimension -get_offset_type: Callable[[OffsetProviderType, str], OffsetProviderTypeElem] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and OffsetProviderType overlap +get_offset_type: Callable[[TableTypes, str], NeighborTableType] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and TableTypes overlap def connectivity_key_over( - offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension | Tag + offset_provider: OffsetProvider | TableTypes, local_dim: Dimension | Tag ) -> str: """ The key of a bound connectivity whose local dimension is `local_dim` (a dimension or its tag). @@ -1554,7 +1600,7 @@ def connectivity_key_over( tag, i.e. its owner's, if bound. Otherwise it is one of the connectivities *sharing* the local dimension (see `NeighborConnectivity`), each keyed by its own tag; the smallest key is taken, so the choice does not depend on the order of the provider. Connectivities sharing a local - dimension have the same neighbor structure (see `check_offset_provider`), so which one does + dimension have the same neighbor structure (see `NeighborConnectivity`), so which one does not matter. Raises: @@ -1578,14 +1624,14 @@ def connectivity_key_over( def _neighbor_dim_of(connectivity: Any) -> Optional[Dimension]: - if isinstance(connectivity, NeighborConnectivityType): - return connectivity.neighbor_dim + if isinstance(connectivity, NeighborTableType): + return connectivity.domain[1] if is_neighbor_table(connectivity): return connectivity.domain.dims[1] return None -def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: str) -> bool: +def has_offset(offset_provider: OffsetProvider | TableTypes, offset_tag: str) -> bool: """Determine if offset provider has an element for the given offset tag.""" try: get_offset(offset_provider, offset_tag) # type: ignore[arg-type] # implementation is shared with `get_offset_type`, no need to duplicate the function @@ -2020,6 +2066,10 @@ class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): #: The connectivity this dimension is the local axis of, or `None` if it indexes no table. #: Set by `NeighborConnectivity` when the connectivity is declared. owner: ClassVar[Optional[type[NeighborConnectivity]]] = None + #: The connectivities sharing this dimension with its owner, in declaration order. Kept so that + #: a table bound under a sharer's `offset_tag` can be typed by its declaration: the table alone + #: looks like the owner's (see `NeighborTableType`). + sharers: ClassVar[tuple[type[NeighborConnectivity], ...]] = () #: Number of entries per element, i.e. the table's second extent, if declared. max_neighbors: ClassVar[Optional[int]] = None #: Least number of *valid* neighbors of any element, if declared. Fewer than @@ -2044,6 +2094,7 @@ def __init_subclass__( # NOTE: reset rather than inherited: a subclass of an owned local dimension is a # different dimension, and does not index its parent's table. cls.owner = None + cls.sharers = () cls.declared_size = _check_neighbor_count(cls, "size", size) cls.max_neighbors = cls.min_neighbors = cls.declared_size @@ -2085,7 +2136,7 @@ class ConnectivityMeta(type): # annotated `Local` makes every declaration's nested class a *variable* for the checkers, so # `Field[Dims[V, V2E.Local]]` is rejected (pyright) or "not valid as a type" (mypy, for the # assigned form). Library code reads it through `local_dimension_of`. - origin: Dimension + domain: Dimension codomain: Dimension @property @@ -2111,7 +2162,7 @@ def _local(cls) -> type[LocalDimensionIndex]: if (local := cls.__dict__.get("Local")) is None: raise TypeError( f"'{cls.__qualname__}' is not a connectivity declaration; declare one by" - " subclassing 'NeighborConnectivity[Origin, Codomain]'." + " subclassing 'NeighborConnectivity[Domain, Codomain]'." ) return cast(type[LocalDimensionIndex], local) @@ -2156,17 +2207,17 @@ def __gt_field_offset__(cls) -> Any: if (field_offset := cls.__dict__.get("_field_offset")) is None: field_offset = fbuiltins.FieldOffset( - cls.offset_tag, source=cls.codomain, target=(cls.origin, cls._local()) + cls.offset_tag, source=cls.codomain, target=(cls.domain, cls._local()) ) type.__setattr__(cls, "_field_offset", field_offset) return field_offset -class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( +class NeighborConnectivity[Domain: DimensionIndex, Codomain: DimensionIndex]( metaclass=ConnectivityMeta ): """ - Declare a neighbor connectivity: for each `Origin` element, a list of `Codomain` neighbors. + Declare a neighbor connectivity: for each `Domain` element, a list of `Codomain` neighbors. The declaration names the connectivity's local dimension -- its nested `Local` class -- and optionally its neighbor counts. It holds no data: the neighbor table is bound at call @@ -2178,7 +2229,7 @@ class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( >>> class Edge(DimensionIndex): ... >>> class V2E(NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): ... class Local(LocalDimensionIndex): ... - >>> V2E.origin is Vertex, V2E.codomain is Edge + >>> V2E.domain is Vertex, V2E.codomain is Edge (True, True) >>> V2E.Local.owner is V2E, V2E.Local.max_neighbors, V2E.Local.min_neighbors (True, 6, 5) @@ -2186,7 +2237,7 @@ class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( # NOTE: `Local` is not annotated (see `ConnectivityMeta`); every subclass declares it, as a # nested class or as `Local: TypeAlias = `. - origin: ClassVar[Dimension] + domain: ClassVar[Dimension] codomain: ClassVar[Dimension] def __init_subclass__( @@ -2211,11 +2262,11 @@ def __init_subclass__( ] if len(params) != 1 or len(params[0]) != 2: raise TypeError( - f"'{name}' must derive from 'NeighborConnectivity[Origin, Codomain]' directly," + f"'{name}' must derive from 'NeighborConnectivity[Domain, Codomain]' directly," " with both dimensions given." ) - origin, codomain = params[0] - for role, dim in (("Origin", origin), ("Codomain", codomain)): + domain, codomain = params[0] + for role, dim in (("Domain", domain), ("Codomain", codomain)): if not isinstance(dim, DimensionMeta) or dim.kind is DimensionKind.LOCAL: raise TypeError(f"'{name}': '{role}' must be a non-local dimension, got '{dim}'.") @@ -2241,11 +2292,11 @@ def __init_subclass__( # Sharing another connectivity's local dimension: the neighbor structure is the # owner's, including its counts, and the sharing connectivity is named by its own tag. owner_name = local.owner.__qualname__ - if origin is not local.owner.origin: + if domain is not local.owner.domain: raise TypeError( - f"'{name}' cannot share the local dimension of '{owner_name}': it has origin" - f" '{origin}', but the neighbors of '{owner_name}' are those of" - f" '{local.owner.origin}'." + f"'{name}' cannot share the local dimension of '{owner_name}': it has domain" + f" '{domain}', but the neighbors of '{owner_name}' are those of" + f" '{local.owner.domain}'." ) for count_name, count in ( ("max_neighbors", max_neighbors), @@ -2257,7 +2308,8 @@ def __init_subclass__( f" shares with '{owner_name}', which declares" f" {count_name}={getattr(local, count_name)}." ) - cls.origin, cls.codomain = origin, codomain + cls.domain, cls.codomain = domain, codomain + local.sharers = (*local.sharers, cls) return for count_name, count in ( ("max_neighbors", max_neighbors), @@ -2286,7 +2338,7 @@ def __init_subclass__( f" ({max_neighbors})." ) - cls.origin, cls.codomain = origin, codomain + cls.domain, cls.codomain = domain, codomain local.owner = cls local.max_neighbors, local.min_neighbors = max_neighbors, min_neighbors @@ -2307,8 +2359,8 @@ def local_dimension_of(connectivity: type[NeighborConnectivity]) -> type[LocalDi def check_neighbor_table( connectivity: type[NeighborConnectivity], - table: NeighborTable | NeighborConnectivityType, -) -> None: + table: NeighborTable | NeighborTableType, +) -> NeighborTableType: """ Check that a neighbor table matches the connectivity declaration it is bound to. @@ -2319,19 +2371,29 @@ def check_neighbor_table( connectivity: The declaration. table: The bound table, or its type (which is all an ahead-of-time compilation has). + Returns: + The type of the table bound to `connectivity`. + Raises: ValueError: On the first mismatch, naming the connectivity and the mismatch. """ - table_type = table if isinstance(table, NeighborConnectivityType) else table.__gt_type__() name = connectivity.__qualname__ local = local_dimension_of(connectivity) def fail(reason: str) -> NoReturn: raise ValueError(f"The table bound to '{name}' does not match its declaration: {reason}.") - if not isinstance(table_type, NeighborConnectivityType): - fail(f"expected a neighbor table, got '{table_type}'") - expected_domain = (connectivity.origin, local) + if isinstance(table, NeighborTableType): + table_type = table + elif is_neighbor_table(table): + table_type = _unbound_table_type(table) + else: + fail(f"expected a neighbor table, got '{table}'") + if isinstance(table_type.connectivity, ConnectivityMeta) and ( + table_type.connectivity is not connectivity + ): + fail(f"its type is bound to '{table_type.connectivity.__qualname__}'") + expected_domain = (connectivity.domain, local) if tuple(table_type.domain) != expected_domain: fail( f"its domain is '({', '.join(map(str, table_type.domain))})'," @@ -2363,3 +2425,4 @@ def fail(reason: str) -> NoReturn: f"min_neighbors == max_neighbors == {max_neighbors} means every element has all" f" its neighbors, but the table has skip value {table_type.skip_value}" ) + return dataclasses.replace(table_type, connectivity=connectivity) diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index 1428e664d1..a464de9763 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -159,9 +159,9 @@ def _make_compiled_programs_pool( def compile( self, - offset_provider: common.OffsetProviderType + offset_provider: common.TableTypes | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + | list[common.TableTypes | common.OffsetProvider] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -185,9 +185,7 @@ def compile( ) if self.compilation_options.connectivities is None and offset_provider is None: - raise ValueError( - "Cannot compile a program without connectivities / OffsetProviderType." - ) + raise ValueError("Cannot compile a program without connectivities / TableTypes.") if not all(isinstance(v, list) for v in static_args.values()): raise TypeError( "Please provide the static arguments as lists." @@ -200,8 +198,7 @@ def compile( offset_provider = [offset_provider] # type: ignore[list-item] # cleanup offset_provider vs offset_provider_type assert all( - common.is_offset_provider(op) or common.is_offset_provider_type(op) - for op in offset_provider + common.is_offset_provider(op) or common.is_table_types(op) for op in offset_provider ) self._compiled_programs.compile(offset_providers=offset_provider, **static_args) @@ -487,9 +484,9 @@ def __call__( @override def compile( self, - offset_provider: common.OffsetProviderType + offset_provider: common.TableTypes | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + | list[common.TableTypes | common.OffsetProvider] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 942a5dcc10..5c9ab763bd 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -130,9 +130,9 @@ def _type_conversion_helper(t: type) -> type[ts.TypeSpec] | tuple[type[ts.TypeSp elif t is common.Dimension: return ts.DimensionType elif t is FieldOffset: - return ts.OffsetType + return ts.ShiftType elif t is common.Connectivity: - return ts.OffsetType + return ts.ShiftType elif t is core_defs.ScalarT: return ts.ScalarType elif t is common.Domain: @@ -493,8 +493,8 @@ def __post_init__(self) -> None: if len(self.target) == 2 and self.target[1].kind != common.DimensionKind.LOCAL: raise ValueError("Second dimension in offset must be a local dimension.") - def __gt_type__(self) -> ts.OffsetType: - return ts.OffsetType(source=self.source, target=self.target, tag=self.value) + def __gt_type__(self) -> ts.ShiftType: + return ts.ShiftType(codomain=self.source, domain=self.target, tag=self.value) @property def Local(self) -> common.Dimension: @@ -541,10 +541,11 @@ def as_connectivity_field(self) -> common.Connectivity: return connectivity -def is_cartesian_offset(offset: FieldOffset | ts.OffsetType) -> bool: +def is_cartesian_offset(offset: FieldOffset | ts.ShiftType) -> bool: + shift_type = offset.__gt_type__() if isinstance(offset, FieldOffset) else offset return ( - len(offset.target) == 1 - and offset.source == offset.target[0] - and offset.source.kind == offset.target[0].kind - and offset.target[0].kind != common.DimensionKind.LOCAL + len(shift_type.domain) == 1 + and shift_type.codomain == shift_type.domain[0] + and shift_type.codomain.kind == shift_type.domain[0].kind + and shift_type.domain[0].kind != common.DimensionKind.LOCAL ) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index ba62641dc5..a486a0765b 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -436,16 +436,19 @@ def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> foast.Attribu new_value = self.visit(node.value, **kwargs) match new_value.type: # `V2E.Local`: the local dimension of a connectivity declaration, which is the last - # target of the offset it is typed as. - case ts.OffsetType(target=(_, local)) if node.attr == "Local": + # dimension of the domain of the shift it is typed as. + case ts.ShiftType(domain=(_, local)) if node.attr == "Local": attr_type: ts.TypeSpec = ts.DimensionType(dim=local) case _: - try: - attr_type = getattr(new_value.type, node.attr) - except AttributeError: + # NOTE: only attributes that are types themselves: a type's other fields (the + # dimensions of a `ShiftType`, say) are not values in DSL code. + if not isinstance( + type_attr := getattr(new_value.type, node.attr, None), ts.TypeSpec + ): raise errors.DSLError( node.location, f"'{new_value.type}' has no attribute '{node.attr}'." - ) from None + ) + attr_type = type_attr return foast.Attribute( value=new_value, attr=node.attr, location=node.location, type=attr_type ) @@ -465,20 +468,21 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri f"Tuples need to be indexed with literal integers, got '{node.index}'.", ) from ex new_type = types[index] - case ts.OffsetType(source=source, target=(target1, target2), tag=tag): - if not target2.kind == DimensionKind.LOCAL: + case ts.ShiftType(codomain=codomain, domain=(domain, local), tag=tag): + if not local.kind == DimensionKind.LOCAL: raise errors.DSLError( new_value.location, "Second dimension in offset must be a local dimension." ) - new_type = ts.OffsetType(source=source, target=(target1,), tag=tag) - case ts.OffsetType(source=source, target=(target,), tag=tag): + new_type = ts.ShiftType(codomain=codomain, domain=(domain,), tag=tag) + case ts.ShiftType(codomain=codomain, domain=(domain,), tag=tag): # for cartesian axes (e.g. I, J) the index of the subscript only # signifies the displacement in the respective dimension, - # but does not change the target type. - if source != target: + # but does not change the domain. + if codomain != domain: raise errors.DSLError( new_value.location, - "Source and target must be equal for offsets with a single target.", + "Codomain and domain must be equal for a shift with a single domain" + " dimension.", ) if tag is None: raise errors.DSLError( @@ -492,7 +496,7 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri ) ], hints=[ - f"Write the displacement directly, e.g. '{source.__qualname__} + 1'." + f"Write the displacement directly, e.g. '{codomain.__qualname__} + 1'." ], ) new_type = new_value.type @@ -634,7 +638,7 @@ def _deduce_compare_type( def _deduce_binop_type( self, node: foast.BinOp, *, left: foast.Expr, right: foast.Expr, **kwargs: Any ) -> Optional[ts.TypeSpec]: - if isinstance(left.type, ts.OffsetType): + if isinstance(left.type, ts.ShiftType): raise errors.DSLError( node.location, f"Type '{left.type}' can not be used in operator '{node.op}'." ) @@ -742,7 +746,7 @@ def _deduce_binop_type( ], ) conn = common.connectivity_for_cartesian_shift(left.type.dim, offset_index) - return ts.OffsetType(source=conn.codomain, target=(conn.domain_dim,)) + return ts.ShiftType(codomain=conn.codomain, domain=(conn.domain_dim,)) else: raise errors.DSLError(node.location, err_msg) @@ -811,8 +815,8 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call: # meaningful unsubscripted, as the neighbor access `field(Off)`. if ( isinstance(arg, (foast.Name, foast.Attribute)) - and isinstance(arg.type, ts.OffsetType) - and len(arg.type.target) == 1 + and isinstance(arg.type, ts.ShiftType) + and len(arg.type.domain) == 1 ): raise errors.DSLError( arg.location, @@ -1009,15 +1013,15 @@ def _visit_astype(self, node: foast.Call, **kwargs: Any) -> foast.Call: def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: arg_0 = node.args[0].type arg_1 = node.args[1].type - assert isinstance(arg_0, ts.OffsetType) + assert isinstance(arg_0, ts.ShiftType) assert isinstance(arg_1, ts.FieldType) if not fbuiltins.is_cartesian_offset(arg_0): - target_dims = ", ".join(d.__qualname__ for d in arg_0.target) # for the diagnostic + domain_dims = ", ".join(d.__qualname__ for d in arg_0.domain) # for the diagnostic raise errors.DSLError( node.location, f"'as_offset' is only supported for Cartesian offsets " - f"(single target dimension equal to source dimension); " - f"got source '{arg_0.source.__qualname__}' and target ({target_dims}).", + f"(a single domain dimension equal to the codomain); " + f"got codomain '{arg_0.codomain.__qualname__}' and domain ({domain_dims}).", ) if not type_info.is_integral(arg_1): raise errors.DSLError( @@ -1027,11 +1031,11 @@ def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: f"{node.location}", ) - if arg_0.source not in arg_1.dims: + if arg_0.codomain not in arg_1.dims: raise errors.DSLError( node.location, f"Incompatible argument in call to '{node.func!s}': " - f"'{arg_0.source}' not in list of offset field dimensions '{arg_1.dims}'. " + f"'{arg_0.codomain}' not in list of offset field dimensions '{arg_1.dims}'. " f"{node.location}", ) diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index f4c8a4fb10..aa6731e11c 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -297,7 +297,7 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: # `field(Off[idx])` # (matched on the type, not the node, to also accept `mod.Off[idx]`) case foast.Subscript( - value=foast.LocatedNode(type=ts.OffsetType(tag=str() as offset_tag)), + value=foast.LocatedNode(type=ts.ShiftType(tag=str() as offset_tag)), index=index, ): # Constant folding to a `Literal` ensures that `index` becomes an `OffsetLiteral`, @@ -331,8 +331,8 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: case foast.Call(func=foast.Name(id="as_offset")): func_args = arg offset_type = func_args.args[0].type - assert isinstance(offset_type, ts.OffsetType) - dim = offset_type.source + assert isinstance(offset_type, ts.ShiftType) + dim = offset_type.codomain offset_field = self.visit(func_args.args[1], **kwargs) current_expr = im.as_fieldop( im.lambda_("__it", "__offset")( @@ -342,7 +342,7 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: ) )(current_expr, offset_field) # `field(Off)` - case foast.LocatedNode(type=ts.OffsetType(tag=str() as offset_tag, target=(_, _))): + case foast.LocatedNode(type=ts.ShiftType(tag=str() as offset_tag, domain=(_, _))): # only a single unstructured shift is supported so returning here is fine even though we # are in a loop. assert len(node.args) == 1 diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index ce78404868..b7cac93ade 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -239,7 +239,7 @@ def _instance_state(obj: Any) -> tuple[Any, ...]: # class, so the strict fingerprinter still checks *its* importability (which is the # connectivity's own, for a nested `Local`, and another module's for a shared one). Deconstruction.from_pieces( - obj.origin, + obj.domain, obj.codomain, common.local_dimension_of(obj), common.local_dimension_of(obj).max_neighbors, diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 3adb6922aa..c54feb6eb8 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -120,15 +120,6 @@ def __init__( def __gt_origin__(self) -> typing.Never: raise NotImplementedError - def __gt_type__(self) -> common.NeighborConnectivityType: - return common.NeighborConnectivityType( - domain=self.domain_dims, - codomain=self.codomain_dim, - max_neighbors=self._max_neighbors, - skip_value=self.skip_value, - dtype=self.dtype, - ) - @property def domain(self) -> common.Domain: return common.Domain( @@ -576,7 +567,7 @@ def execute_shift( common.connectivity_key_over(offset_provider, tag), ) assert common.is_neighbor_table(offset_implementation) - source_dim = offset_implementation.__gt_type__().source_dim + source_dim = offset_implementation.__gt_type__().domain[0] cur_index = pos[source_dim.tag] assert common.is_int_index(cur_index) if offset_implementation[cur_index, index].as_scalar() in [ @@ -598,7 +589,7 @@ def execute_shift( return new_pos offset_implementation = common.get_offset(offset_provider, tag) if common.is_neighbor_table(offset_implementation): - source_dim = offset_implementation.__gt_type__().source_dim + source_dim = offset_implementation.__gt_type__().domain[0] assert source_dim.tag in pos new_pos = pos.copy() new_pos.pop(source_dim.tag) @@ -1436,7 +1427,7 @@ def local_dim(self) -> common.Dimension: assert offset_provider is not None connectivity = common.get_offset(offset_provider, offset_tag) assert common.is_neighbor_table(connectivity) - return connectivity.__gt_type__().neighbor_dim + return connectivity.__gt_type__().domain[1] @dataclasses.dataclass(frozen=True) @@ -1466,7 +1457,7 @@ def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: return _List( values=tuple( shifted.deref() - for i in range(connectivity.__gt_type__().max_neighbors) + for i in range(len(connectivity.domain[1].unit_range)) if (shifted := it.shift(offset_str, i)).can_deref() ), offset=offset, @@ -1549,7 +1540,7 @@ def deref(self) -> Any: return _List( values=tuple( shifted.deref() - for i in range(connectivity.__gt_type__().max_neighbors) + for i in range(len(connectivity.domain[1].unit_range)) if ( shifted := self.it.shift(*self.offsets, SparseTag(self.list_offset), i) ).can_deref() @@ -1685,9 +1676,9 @@ def _dimension_to_tag( return {k.tag: v for k, v in domain.items()} -def _validate_domain(domain: Domain, offset_provider_type: common.OffsetProviderType) -> None: +def _validate_domain(domain: Domain, offset_provider_type: common.TableTypes) -> None: if isinstance(domain, runtime.CartesianDomain): - if any(isinstance(o, common.ConnectivityType) for o in offset_provider_type.values()): + if any(isinstance(o, common.NeighborTableType) for o in offset_provider_type.values()): raise RuntimeError( "Got a 'CartesianDomain', but found a 'Connectivity' in 'offset_provider', expected 'UnstructuredDomain'." ) @@ -1807,7 +1798,7 @@ def _fieldspec_list_to_value( assert common.is_neighbor_table(connectivity) return domain.insert( len(domain), - common.named_range((offset_type, connectivity.__gt_type__().max_neighbors)), + common.named_range((offset_type, len(connectivity.domain[1].unit_range))), ), type_.element_type return domain, type_ diff --git a/src/gt4py/next/iterator/ir_utils/domain_utils.py b/src/gt4py/next/iterator/ir_utils/domain_utils.py index b23ef3a934..b87f340096 100644 --- a/src/gt4py/next/iterator/ir_utils/domain_utils.py +++ b/src/gt4py/next/iterator/ir_utils/domain_utils.py @@ -172,7 +172,7 @@ def translate( | Literal[trace_shifts.Sentinel.VALUE, trace_shifts.Sentinel.ALL_NEIGHBORS], ..., ], - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, #: A dictionary mapping axes names to their length. See #: func:`gt4py.next.iterator.transforms.infer_domain.infer_expr` for more details. symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, @@ -208,13 +208,13 @@ def translate( trace_shifts.Sentinel.VALUE, ] - connectivity: common.NeighborTable | common.NeighborConnectivityType + connectivity: common.NeighborTable | common.NeighborTableType if common.is_offset_provider(offset_provider): connectivity = common.get_offset(offset_provider, off.value) old_dim = connectivity.domain.dims[0] new_dim = connectivity.codomain else: - assert common.is_offset_provider_type(offset_provider) + assert common.is_table_types(offset_provider) connectivity = common.get_offset_type(offset_provider, off.value) old_dim = connectivity.domain[0] new_dim = connectivity.codomain diff --git a/src/gt4py/next/iterator/runtime.py b/src/gt4py/next/iterator/runtime.py index 88a466229f..3b604fcb31 100644 --- a/src/gt4py/next/iterator/runtime.py +++ b/src/gt4py/next/iterator/runtime.py @@ -133,9 +133,7 @@ def fendef( ) -def _deduce_domain( - domain: dict[common.Dimension, range], offset_provider_type: common.OffsetProviderType -): +def _deduce_domain(domain: dict[common.Dimension, range], offset_provider_type: common.TableTypes): if isinstance(domain, UnstructuredDomain): domain_builtin = builtins.unstructured_domain elif isinstance(domain, CartesianDomain): @@ -143,7 +141,7 @@ def _deduce_domain( else: domain_builtin = ( builtins.unstructured_domain - if any(isinstance(o, common.ConnectivityType) for o in offset_provider_type.values()) + if any(isinstance(o, common.NeighborTableType) for o in offset_provider_type.values()) else builtins.cartesian_domain ) diff --git a/src/gt4py/next/iterator/transforms/collapse_tuple.py b/src/gt4py/next/iterator/transforms/collapse_tuple.py index 08ad5c9ede..1d049897d9 100644 --- a/src/gt4py/next/iterator/transforms/collapse_tuple.py +++ b/src/gt4py/next/iterator/transforms/collapse_tuple.py @@ -187,7 +187,7 @@ def apply( node: itir.Node, *, remove_letified_make_tuple_elements: bool = True, - offset_provider_type: Optional[common.OffsetProviderType] = None, + offset_provider_type: Optional[common.TableTypes] = None, within_stencil: Optional[bool] = None, # manually passing enabled transformations is mostly for allowing separate testing of the modes enabled_transformations: Optional[Transformation] = None, diff --git a/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py b/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py index 40d956fca0..aaf34d5759 100644 --- a/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py +++ b/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py @@ -27,7 +27,7 @@ def apply( cls, node: itir.Node, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, allow_undeclared_symbols: bool = False, ) -> itir.Node: node = type_inference.infer( diff --git a/src/gt4py/next/iterator/transforms/cse.py b/src/gt4py/next/iterator/transforms/cse.py index 983fcafbd1..78d4612a16 100644 --- a/src/gt4py/next/iterator/transforms/cse.py +++ b/src/gt4py/next/iterator/transforms/cse.py @@ -462,7 +462,7 @@ def apply( cls, node: ProgramOrExpr, within_stencil: bool | None = None, - offset_provider_type: common.OffsetProviderType | None = None, + offset_provider_type: common.TableTypes | None = None, *, uids: utils.IDGeneratorPool, ) -> ProgramOrExpr: diff --git a/src/gt4py/next/iterator/transforms/dead_code_elimination.py b/src/gt4py/next/iterator/transforms/dead_code_elimination.py index 1ea906ae98..8a7d84b0f2 100644 --- a/src/gt4py/next/iterator/transforms/dead_code_elimination.py +++ b/src/gt4py/next/iterator/transforms/dead_code_elimination.py @@ -17,7 +17,7 @@ def dead_code_elimination( program: itir.Program, *, uids: utils.IDGeneratorPool, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> itir.Program: """ Perform dead code elimination on a program by simplifying or removing diff --git a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py index 6b705a29dc..7028e0837e 100644 --- a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py +++ b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py @@ -61,7 +61,7 @@ def apply( node: ProgramOrExpr, *, uids: utils.IDGeneratorPool | None, - offset_provider_type: common.OffsetProviderType | None = None, + offset_provider_type: common.TableTypes | None = None, ) -> ProgramOrExpr: if node.type is None: node = itir_inference.infer( diff --git a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py index 3e71508e25..9babe30fe8 100644 --- a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py +++ b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py @@ -123,7 +123,7 @@ def fuse_as_fieldop( expr: itir.Expr, eligible_args: list[bool], *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, enable_cse: bool, uids: utils.IDGeneratorPool, ) -> itir.Expr: @@ -301,7 +301,7 @@ def all(self) -> FuseAsFieldOp.Transformation: enabled_transformations = Transformation.all() uids: utils.IDGeneratorPool - offset_provider_type: common.OffsetProviderType + offset_provider_type: common.TableTypes enable_cse: bool # option to disable is mainly for testing purposes @classmethod @@ -309,7 +309,7 @@ def apply( cls, node: itir.Program, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, uids: utils.IDGeneratorPool, allow_undeclared_symbols=False, within_set_at_expr: Optional[bool] = None, diff --git a/src/gt4py/next/iterator/transforms/global_tmps.py b/src/gt4py/next/iterator/transforms/global_tmps.py index b2eafc1090..8952554479 100644 --- a/src/gt4py/next/iterator/transforms/global_tmps.py +++ b/src/gt4py/next/iterator/transforms/global_tmps.py @@ -311,7 +311,7 @@ def _transform_stmt( def create_global_tmps( program: itir.Program, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, #: A dictionary mapping axes names to their length. See :func:`infer_domain.infer_expr` for #: more details. symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, diff --git a/src/gt4py/next/iterator/transforms/infer_domain.py b/src/gt4py/next/iterator/transforms/infer_domain.py index 5466843476..8df09b0860 100644 --- a/src/gt4py/next/iterator/transforms/infer_domain.py +++ b/src/gt4py/next/iterator/transforms/infer_domain.py @@ -58,7 +58,7 @@ class DomainAccessDescriptor(eve.StrEnum): class InferenceOptions(typing.TypedDict): - offset_provider: common.OffsetProvider | common.OffsetProviderType + offset_provider: common.OffsetProvider | common.TableTypes symbolic_domain_sizes: dict[str, itir.Expr] | None allow_uninferred: bool keep_existing_domains: bool @@ -130,7 +130,7 @@ def _extract_accessed_domains( stencil: itir.Expr, input_ids: list[str], target_domain: NonTupleDomainAccess, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]], ) -> dict[str, NonTupleDomainAccess]: accessed_domains: dict[str, NonTupleDomainAccess] = {} @@ -186,7 +186,7 @@ def _infer_as_fieldop( applied_fieldop: itir.FunCall, target_domain: DomainAccess, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]], allow_uninferred: bool, keep_existing_domains: bool, @@ -445,7 +445,7 @@ def infer_expr( expr: _Expr_T, domain: DomainAccess, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, allow_uninferred: bool = False, keep_existing_domains: bool = False, @@ -573,7 +573,7 @@ def _infer_stmt( def infer_program( program: itir.Program, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, allow_uninferred: bool = False, keep_existing_domains: bool = False, diff --git a/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py b/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py index 66cc3af85a..489d140027 100644 --- a/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py +++ b/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py @@ -33,14 +33,14 @@ def _dynamic_shift_args(node: itir.Expr) -> list[bool] | None: @dataclasses.dataclass class InlineDynamicShifts(eve.NodeTranslator, eve.VisitorWithSymbolTableTrait): - offset_provider_type: common.OffsetProviderType + offset_provider_type: common.TableTypes uids: utils.IDGeneratorPool @classmethod def apply( cls, node: itir.Program, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, uids: utils.IDGeneratorPool, ): return cls(offset_provider_type=offset_provider_type, uids=uids).visit(node) diff --git a/src/gt4py/next/iterator/transforms/inline_scalar.py b/src/gt4py/next/iterator/transforms/inline_scalar.py index b424074b5c..223a484702 100644 --- a/src/gt4py/next/iterator/transforms/inline_scalar.py +++ b/src/gt4py/next/iterator/transforms/inline_scalar.py @@ -19,7 +19,7 @@ class InlineScalar(eve.NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) @classmethod - def apply(cls, program: itir.Program, offset_provider_type: common.OffsetProviderType): + def apply(cls, program: itir.Program, offset_provider_type: common.TableTypes): program = itir_inference.infer(program, offset_provider_type=offset_provider_type) return cls().visit(program) diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index ec4363aba3..0865ffda00 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -55,7 +55,7 @@ def _max_domain_range_sizes(offset_provider: common.OffsetProvider) -> dict[str, sizes: dict[str, int] = {} for provider in offset_provider.values(): if common.is_neighbor_table(provider): - src_dim = provider.__gt_type__().source_dim.tag + src_dim = provider.__gt_type__().domain[0].tag codomain_dim = provider.__gt_type__().codomain.tag sizes[src_dim] = max(sizes.get(src_dim, 0), provider.ndarray.shape[0]) sizes[codomain_dim] = max( @@ -134,7 +134,7 @@ def _process_symbolic_domains_option( def apply_common_transforms( ir: itir.Program, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, extract_temporaries=False, unroll_reduce=False, common_subexpression_elimination=True, @@ -147,7 +147,7 @@ def apply_common_transforms( use_max_domain_range_on_unstructured_shift: Optional[bool] = None, ) -> itir.Program: assert isinstance(ir, itir.Program) - # TODO(tehrengruber): Allow `common.OffsetProviderType`, but domain inference currently + # TODO(tehrengruber): Allow `common.TableTypes`, but domain inference currently # relies on static information or `symbolic_domain_sizes`. assert common.is_offset_provider(offset_provider) diff --git a/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py index b9ade1d636..4dd3f598f8 100644 --- a/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py +++ b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py @@ -35,7 +35,7 @@ def _broadcast_to(expr: itir.Expr, target_dims: list[common.Dimension]) -> itir. def _concat_where_with_explicit_broadcast( node: itir.FunCall, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, ) -> itir.FunCall: """ @@ -107,7 +107,7 @@ class _PruneEmptyConcatWhere(PreserveLocationVisitor, NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) - offset_provider: common.OffsetProvider | common.OffsetProviderType + offset_provider: common.OffsetProvider | common.TableTypes symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None @classmethod @@ -115,7 +115,7 @@ def apply( cls: type[Self], node: PRG, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, ) -> PRG: return cls( diff --git a/src/gt4py/next/iterator/transforms/unroll_reduce.py b/src/gt4py/next/iterator/transforms/unroll_reduce.py index 22491346bb..79067bf663 100644 --- a/src/gt4py/next/iterator/transforms/unroll_reduce.py +++ b/src/gt4py/next/iterator/transforms/unroll_reduce.py @@ -52,18 +52,18 @@ def _get_partial_local_dims(reduce_args: Iterable[itir.Expr]) -> Iterable[common def _get_connectivity( applied_reduce_node: itir.FunCall, - offset_provider_type: common.OffsetProviderType, -) -> common.NeighborConnectivityType: + offset_provider_type: common.TableTypes, +) -> common.NeighborTableType: """Return single connectivity that is compatible with the arguments of the reduce.""" if not cpm.is_applied_reduce(applied_reduce_node): raise ValueError("Expected a call to a 'reduce' object, i.e. 'reduce(...)(...)'.") - connectivities: list[common.NeighborConnectivityType] = [] + connectivities: list[common.NeighborTableType] = [] for local_dim in _get_partial_local_dims(applied_reduce_node.args): conn = common.get_offset_type( offset_provider_type, common.connectivity_key_over(offset_provider_type, local_dim) ) - assert isinstance(conn, common.NeighborConnectivityType) + assert isinstance(conn, common.NeighborTableType) connectivities.append(conn) if not connectivities: @@ -87,13 +87,13 @@ class UnrollReduce(PreserveLocationVisitor, NodeTranslator): def apply( cls, node: itir.Node, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, uids: utils.IDGeneratorPool, ) -> itir.Node: return cls(uids=uids).visit(node, offset_provider_type=offset_provider_type) def _visit_reduce( - self, node: itir.FunCall, offset_provider_type: common.OffsetProviderType + self, node: itir.FunCall, offset_provider_type: common.TableTypes ) -> itir.Expr: connectivity_type = _get_connectivity(node, offset_provider_type) max_neighbors = connectivity_type.max_neighbors diff --git a/src/gt4py/next/iterator/type_system/inference.py b/src/gt4py/next/iterator/type_system/inference.py index b878640f12..77d0b033a8 100644 --- a/src/gt4py/next/iterator/type_system/inference.py +++ b/src/gt4py/next/iterator/type_system/inference.py @@ -182,7 +182,7 @@ def on_type_ready(self, cb: Callable[[ts.TypeSpec], None]) -> None: def __call__( self, *args: type_synthesizer.TypeOrTypeSynthesizer, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, **kwargs, ) -> Union[ts.TypeSpec, ObservableTypeSynthesizer]: assert all(isinstance(arg, (ts.TypeSpec, ObservableTypeSynthesizer)) for arg in args), ( @@ -256,7 +256,7 @@ class ITIRTypeInference(eve.NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) - offset_provider_type: Optional[common.OffsetProviderType] + offset_provider_type: Optional[common.TableTypes] #: Allow sym refs to symbols that have not been declared. Mostly used in testing. allow_undeclared_symbols: bool #: Reinference-mode skipping already typed nodes. @@ -267,7 +267,7 @@ def apply( cls, node: T, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, inplace: bool = False, allow_undeclared_symbols: bool = False, ) -> T: @@ -352,7 +352,7 @@ def apply( @classmethod def apply_reinfer( - cls, node: T, *, offset_provider_type: Optional[common.OffsetProviderType] = None + cls, node: T, *, offset_provider_type: Optional[common.TableTypes] = None ) -> T: """ Given a partially typed node infer the type of ``node`` and its sub-nodes. diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 278683429e..859d8ae61d 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -68,7 +68,7 @@ def __post_init__(self): def __call__( self, *args: TypeOrTypeSynthesizer, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, **kwargs, ) -> TypeOrTypeSynthesizer: return self.type_synthesizer(*args, offset_provider_type=offset_provider_type, **kwargs) @@ -313,22 +313,22 @@ def broadcast( def neighbors( offset_literal: it_ts.OffsetLiteralType, it: it_ts.IteratorType, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> ts.ListType: assert isinstance(offset_literal, it_ts.OffsetLiteralType) and isinstance( offset_literal.value, str ) assert isinstance(it, it_ts.IteratorType) conn_type = common.get_offset_type(offset_provider_type, offset_literal.value) - assert isinstance(conn_type, common.NeighborConnectivityType) - return ts.ListType(element_type=it.element_type, offset_type=conn_type.neighbor_dim) + assert isinstance(conn_type, common.NeighborTableType) + return ts.ListType(element_type=it.element_type, offset_type=conn_type.domain[1]) @_register_builtin_type_synthesizer def lift(stencil: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer def apply_lift( - *its: it_ts.IteratorType, offset_provider_type: common.OffsetProviderType + *its: it_ts.IteratorType, offset_provider_type: common.TableTypes ) -> it_ts.IteratorType: assert all(isinstance(it, it_ts.IteratorType) for it in its) stencil_args = [ @@ -451,7 +451,7 @@ def _canonicalize_nb_fields( def _resolve_dimensions( input_dims: list[common.Dimension], shift_tuple: tuple[itir.OffsetLiteral | itir.CartesianOffset, ...], - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> list[common.Dimension]: """ Resolves the final dimensions by applying shifts from the given shift tuple. @@ -489,21 +489,16 @@ def _resolve_dimensions( ... itir.OffsetLiteral(value="V2E"), ... itir.OffsetLiteral(value=0), ... ) + >>> def table_type(domain, codomain, max_neighbors): # of tables no declaration names + ... structure = common.ConnectivityType( + ... domain=domain, codomain=codomain, skip_value=None, dtype=None + ... ) + ... return common.NeighborTableType( + ... connectivity=structure, dtype=None, skip_value=None, max_neighbors=max_neighbors + ... ) >>> offset_provider_type = { - ... "C2V": common.NeighborConnectivityType( - ... domain=(Cell, C2V), - ... codomain=Vertex, - ... skip_value=None, - ... dtype=None, - ... max_neighbors=3, - ... ), - ... "V2E": common.NeighborConnectivityType( - ... domain=(Vertex, V2E), - ... codomain=Edge, - ... skip_value=None, - ... dtype=None, - ... max_neighbors=4, - ... ), + ... "C2V": table_type((Cell, C2V), Vertex, 3), + ... "V2E": table_type((Vertex, V2E), Edge, 4), ... } >>> _resolve_dimensions(input_dims, shift_tuple, offset_provider_type) [gt4py.next.iterator.type_system.type_synthesizer.Cell[horizontal], gt4py.next.iterator.type_system.type_synthesizer.K[vertical]] @@ -555,7 +550,7 @@ def _resolve_dimensions( off_literal.value, str ) offset_type = common.get_offset_type(offset_provider_type, off_literal.value) - if isinstance(offset_type, common.NeighborConnectivityType): + if isinstance(offset_type, common.NeighborTableType): if resolved_dim == offset_type.codomain: # Check if input fits to offset resolved_dim = offset_type.domain[0] # Update input_dim for next iteration else: @@ -571,7 +566,7 @@ def as_fieldop( stencil: TypeSynthesizer, domain: Optional[ts.DomainType] = None, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> TypeSynthesizer: @type_synthesizer def applied_as_fieldop( @@ -647,7 +642,7 @@ def scan( @type_synthesizer def apply_scan( - *its: it_ts.IteratorType, offset_provider_type: common.OffsetProviderType + *its: it_ts.IteratorType, offset_provider_type: common.TableTypes ) -> ts.DataType: result = scan_pass(init, *its, offset_provider_type=offset_provider_type) assert isinstance(result, ts.DataType) @@ -659,9 +654,7 @@ def apply_scan( @_register_builtin_type_synthesizer def map_list(op: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer - def applied_map( - *args: ts.ListType, offset_provider_type: common.OffsetProviderType - ) -> ts.ListType: + def applied_map(*args: ts.ListType, offset_provider_type: common.TableTypes) -> ts.ListType: assert len(args) > 0 assert all(isinstance(arg, ts.ListType) for arg in args) arg_el_types = [arg.element_type for arg in args] @@ -682,9 +675,7 @@ def _make_tuple_map_synthesizer( def tuple_map_synthesizer(op: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer - def applied_map( - arg: ts.TupleType, offset_provider_type: common.OffsetProviderType - ) -> ts.TupleType: + def applied_map(arg: ts.TupleType, offset_provider_type: common.TableTypes) -> ts.TupleType: if not isinstance(arg, ts.TupleType): raise TypeError( f"'{builtin_name}' requires a 'TupleType' argument, got '{type(arg).__name__}'." @@ -718,7 +709,7 @@ def applied_map( @_register_builtin_type_synthesizer def reduce(op: TypeSynthesizer, init: ts.TypeSpec) -> TypeSynthesizer: @type_synthesizer - def applied_reduce(*args: ts.ListType, offset_provider_type: common.OffsetProviderType): + def applied_reduce(*args: ts.ListType, offset_provider_type: common.TableTypes): assert all(isinstance(arg, ts.ListType) for arg in args) assert any( arg.offset_type is not None for arg in args @@ -731,7 +722,7 @@ def applied_reduce(*args: ts.ListType, offset_provider_type: common.OffsetProvid @_register_builtin_type_synthesizer -def shift(*offset_literals, offset_provider_type: common.OffsetProviderType) -> TypeSynthesizer: +def shift(*offset_literals, offset_provider_type: common.TableTypes) -> TypeSynthesizer: @type_synthesizer def apply_shift( it: it_ts.IteratorType | ts.DeferredType, @@ -754,7 +745,7 @@ def apply_shift( assert isinstance(offset_axis, it_ts.OffsetLiteralType) assert isinstance(offset_axis.value, str) type_ = common.get_offset_type(offset_provider_type, offset_axis.value) - assert isinstance(type_, common.NeighborConnectivityType) + assert isinstance(type_, common.NeighborTableType) source_dim, target_dim = type_.domain[0], type_.codomain found = False diff --git a/src/gt4py/next/otf/arguments.py b/src/gt4py/next/otf/arguments.py index 67c9f2bdc4..c0f86a96f6 100644 --- a/src/gt4py/next/otf/arguments.py +++ b/src/gt4py/next/otf/arguments.py @@ -135,7 +135,7 @@ class CompileTimeArgs: args: tuple[ts.TypeSpec, ...] kwargs: dict[str, ts.TypeSpec] - offset_provider: common.OffsetProvider # TODO(havogt): replace with common.OffsetProviderType once the temporary pass doesn't require the runtime information + offset_provider: common.OffsetProvider # TODO(havogt): replace with common.TableTypes once the temporary pass doesn't require the runtime information column_axis: Optional[common.Dimension] #: A mapping from an argument descriptor type to a context containing the actual descriptors. #: If an argument or element of an argument has no descriptor, the respective value is `None`. @@ -144,7 +144,7 @@ class CompileTimeArgs: argument_descriptor_contexts: ArgStaticDescriptorsContextsByType @property - def offset_provider_type(self) -> common.OffsetProviderType: + def offset_provider_type(self) -> common.TableTypes: return common.offset_provider_to_type(self.offset_provider) @classmethod diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index b71fdd9df9..f614b0956e 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -91,7 +91,7 @@ def compile_variant_hook( key: CompiledProgramsKey, backend: gtx_backend.Backend, argument_descriptors: ArgStaticDescriptorsByType, - offset_provider: common.OffsetProviderType | common.OffsetProvider, + offset_provider: common.TableTypes | common.OffsetProvider, ) -> None: """Callback hook invoked before compiling a program variant.""" @@ -627,7 +627,7 @@ def _finish_compilation_job(self, key: CompiledProgramsKey) -> bool: def _compile_variant( self, argument_descriptors: ArgStaticDescriptorsByType, - offset_provider: common.OffsetProviderType | common.OffsetProvider, + offset_provider: common.TableTypes | common.OffsetProvider, #: tuple consisting of the types of the positional and keyword arguments. arg_specialization_info: tuple[tuple[ts.TypeSpec, ...], dict[str, ts.TypeSpec]] | None = None, @@ -636,9 +636,9 @@ def _compile_variant( call_key: CompiledProgramsKey | None = None, ) -> None: if not common.is_offset_provider(offset_provider): - if common.is_offset_provider_type(offset_provider): + if common.is_table_types(offset_provider): raise ValueError( - "Variant compilation of programs with 'OffsetProviderType' is not yet supported." + "Variant compilation of programs with 'TableTypes' is not yet supported." ) else: raise ValueError(f"Invalid 'offset_provider': {offset_provider}") @@ -709,11 +709,11 @@ def _compile_variant( # domains and of scans. def compile( self, - offset_providers: list[common.OffsetProvider | common.OffsetProviderType], + offset_providers: list[common.OffsetProvider | common.TableTypes], **static_args: list[ScalarOrTupleOfScalars], ) -> None: """ - Compiles the program for all combinations of static arguments and the given 'OffsetProviderType'. + Compiles the program for all combinations of static arguments and the given 'TableTypes'. Note: In case you want to compile for specific combinations of static arguments (instead of the combinatoral), you can call compile multiples times. diff --git a/src/gt4py/next/otf/options.py b/src/gt4py/next/otf/options.py index 4f77d44586..6f0d7ac7a7 100644 --- a/src/gt4py/next/otf/options.py +++ b/src/gt4py/next/otf/options.py @@ -32,7 +32,7 @@ class CompilationOptions: #: when jitting is enabled, or on a call to `compile`. static_params: Sequence[str] | None = None - # TODO(ricoh): replace with common.OffsetProviderType once the temporary pass doesn't require the runtime information + # TODO(ricoh): replace with common.TableTypes once the temporary pass doesn't require the runtime information #: A dictionary holding static/compile-time information about the offset providers. #: For now, it is used for ahead of time compilation in DaCe orchestrated programs, #: i.e. DaCe programs that call GT4Py Programs -SDFGConvertible interface-. diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index 78f4c7190c..45bec4de35 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -68,7 +68,7 @@ def _process_regular_arguments( self, program: itir.Program, arg_types: tuple[ts.TypeSpec, ...], - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> tuple[list[interface.Parameter], list[str]]: parameters: list[interface.Parameter] = [] arg_exprs: list[str] = [] @@ -98,20 +98,20 @@ def _process_regular_arguments( if isinstance(dim, fbuiltins.FieldOffset) else common.connectivity_key_over(offset_provider_type, dim), ) - assert isinstance(connectivity, common.NeighborConnectivityType) + assert isinstance(connectivity, common.NeighborTableType) size = connectivity.max_neighbors arg = f"gridtools::sid::dimension_to_tuple_like({arg})" arg_exprs.append(arg) return parameters, arg_exprs def _process_connectivity_args( - self, offset_provider_type: common.OffsetProviderType + self, offset_provider_type: common.TableTypes ) -> tuple[list[interface.Parameter], list[str]]: parameters: list[interface.Parameter] = [] arg_exprs: list[str] = [] for name, connectivity_type in offset_provider_type.items(): - if isinstance(connectivity_type, common.NeighborConnectivityType): + if isinstance(connectivity_type, common.NeighborTableType): if connectivity_type.dtype.scalar_type not in [np.int32, np.int64]: raise ValueError( "Neighbor table indices must be of type 'np.int32' or 'np.int64'." @@ -147,7 +147,7 @@ def _process_connectivity_args( ) else: raise AssertionError( - f"Expected offset provider type '{name}' to be a 'NeighborConnectivityType', " + f"Expected offset provider type '{name}' to be a 'NeighborTableType', " f"got '{type(connectivity_type).__name__}'." ) @@ -156,7 +156,7 @@ def _process_connectivity_args( def _preprocess_program( self, program: itir.Program, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, ) -> itir.Program: return pass_manager.apply_common_transforms( program, @@ -170,7 +170,7 @@ def _preprocess_program( def generate_stencil_source( self, program: itir.Program, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, column_axis: Optional[common.Dimension], ) -> str: if self.enable_itir_transforms: diff --git a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py index 38fe8fa1c6..c893be66e9 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py @@ -159,7 +159,7 @@ def _collect_dimensions_from_params( def _collect_offset_definitions( node: itir.Node, grid_type: common.GridType, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> dict[str, TagDefinition]: offset_definitions = {} offset_provider_type = {**offset_provider_type} @@ -188,17 +188,17 @@ def _collect_offset_definitions( ) for offset_name, connectivity_type in offset_provider_type.items(): - if isinstance(connectivity_type, common.NeighborConnectivityType): + if isinstance(connectivity_type, common.NeighborTableType): assert grid_type == common.GridType.UNSTRUCTURED offset_definitions[offset_name] = TagDefinition( name=Sym(id=common.codegen_name(offset_name)) ) - if offset_name != connectivity_type.neighbor_dim.tag: - offset_definitions[connectivity_type.neighbor_dim.tag] = TagDefinition( - name=Sym(id=common.codegen_name(connectivity_type.neighbor_dim.tag)) + if offset_name != connectivity_type.domain[1].tag: + offset_definitions[connectivity_type.domain[1].tag] = TagDefinition( + name=Sym(id=common.codegen_name(connectivity_type.domain[1].tag)) ) - for dim in [connectivity_type.source_dim, connectivity_type.codomain]: + for dim in [connectivity_type.domain[0], connectivity_type.codomain]: if dim.kind != common.DimensionKind.HORIZONTAL: raise NotImplementedError() offset_definitions[dim.tag] = TagDefinition( @@ -206,7 +206,7 @@ def _collect_offset_definitions( ) else: raise AssertionError( - "Elements of the offset provider type need to be a 'NeighborConnectivityType'." + "Elements of the offset provider type need to be a 'NeighborTableType'." ) return offset_definitions @@ -339,7 +339,7 @@ class GTFN_lowering(eve.NodeTranslator, eve.VisitorWithSymbolTableTrait): } _unary_op_map: ClassVar[dict[str, str]] = {"not_": "!"} - offset_provider_type: common.OffsetProviderType + offset_provider_type: common.TableTypes column_axis: Optional[common.Dimension] grid_type: common.GridType @@ -354,7 +354,7 @@ def apply( cls, node: itir.Program, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, column_axis: Optional[common.Dimension], ) -> Program: if not isinstance(node, itir.Program): @@ -506,7 +506,7 @@ def _visit_unstructured_domain(self, node: itir.FunCall, **kwargs: Any) -> Node: for o in shift_offsets: if o in self.offset_provider_type and isinstance( common.get_offset_type(self.offset_provider_type, o), - common.NeighborConnectivityType, + common.NeighborTableType, ): # `o` is an offset-provider key, i.e. a qualified tag: mangle it exactly as # its `TagDefinition` was, or the reference names an undeclared tag type. diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index d3026b3ea2..d95968c870 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -87,7 +87,7 @@ class DataflowBuilder(Protocol): """Visitor interface to build a dataflow subgraph.""" @abc.abstractmethod - def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: ... + def get_offset_provider_type(self, offset: str) -> gtx_common.NeighborTableType: ... @abc.abstractmethod def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: @@ -560,13 +560,13 @@ class GTIRToSDFG(eve.NodeVisitor, SDFGBuilder): from where to continue building the SDFG. """ - offset_provider_type: gtx_common.OffsetProviderType + offset_provider_type: gtx_common.TableTypes column_axis: Optional[gtx_common.Dimension] uids: gtx_utils.IDGeneratorPool = dataclasses.field( init=False, repr=False, default_factory=lambda: gtx_utils.IDGeneratorPool() ) - def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: + def get_offset_provider_type(self, offset: str) -> gtx_common.NeighborTableType: return gtx_common.get_offset_type(self.offset_provider_type, offset) def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: @@ -1040,7 +1040,7 @@ def _add_sdfg_params( self.offset_provider_type ).items(): gt_type = ts.FieldType( - dims=[connectivity_type.source_dim, connectivity_type.neighbor_dim], + dims=[connectivity_type.domain[0], connectivity_type.domain[1]], dtype=tt.from_dtype(connectivity_type.dtype), ) # We store all connectivity tables as transient arrays here; later, while building @@ -1396,7 +1396,7 @@ def visit_SymRef( def lower_program_to_sdfg( ir: gtir.Program, - offset_provider_type: gtx_common.OffsetProviderType, + offset_provider_type: gtx_common.TableTypes, column_axis: Optional[gtx_common.Dimension] = None, ) -> dace.SDFG: """ diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py index 60827eeedc..0eb550936c 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py @@ -257,7 +257,7 @@ def translate_concat_where( offset_provider_type = sdfg_builder.get_offset_provider_type( sdfg_builder.connectivity_key_over(local_dim) ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) local_idx = gtx_common.order_dimensions([*output_dims, local_dim]).index(local_dim) output_shape.insert(local_idx, offset_provider_type.max_neighbors) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 9ba625ea26..98e5d7a355 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -765,7 +765,7 @@ def _visit_if_branch_arg( self.subgraph_builder.get_offset_provider_type( self.subgraph_builder.connectivity_key_over(local_dim) ), - gtx_common.NeighborConnectivityType, + gtx_common.NeighborTableType, ) # find position of the local dimension in the field layout assert isinstance(arg_desc, dace.data.Array) @@ -1079,7 +1079,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: offset = node.args[0].value assert isinstance(offset, str) conn_type = self.subgraph_builder.get_offset_provider_type(offset) - assert isinstance(conn_type, gtx_common.NeighborConnectivityType) + assert isinstance(conn_type, gtx_common.NeighborTableType) it = self.visit(node.args[1]) assert isinstance(it, IteratorExpr) @@ -1092,8 +1092,8 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: origin for dim, origin in it.field_domain if dim == conn_type.codomain ) # make sure that the iterator can access the connectivity table - assert conn_type.source_dim in it.indices - conn_source_index = it.indices[conn_type.source_dim] + assert conn_type.domain[0] in it.indices + conn_source_index = it.indices[conn_type.domain[0]] assert isinstance(conn_source_index, SymbolExpr) # initially, the storage for the connectivty tables is created as transient; @@ -1154,8 +1154,8 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: # NOTE: the connectivity's own local dimension, not one synthesized from the offset # tag. The latter named a local dimension after the *offset*, which only coincided with # the real one under the old `V2EDim = Dimension("V2E")` convention, and under nominal - # identity (ADR 0029) a tag string cannot be turned back into a dimension at all. - offset_type = conn_type.neighbor_dim + # identity (ADR 0028) a tag string cannot be turned back into a dimension at all. + offset_type = conn_type.domain[1] neighbor_idx = gtir_to_sdfg_utils.get_map_variable(offset_type) index_connector = "__index" @@ -1309,7 +1309,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: tasklet_expression = f"{output_connector} = {fun_python_code}" input_args = [self.visit(arg) for arg in node.args] - input_conn_types: dict[gtx_common.Dimension, gtx_common.NeighborConnectivityType] = {} + input_conn_types: dict[gtx_common.Dimension, gtx_common.NeighborTableType] = {} for input_arg in input_args: assert isinstance(input_arg.gt_dtype, ts.ListType) assert input_arg.gt_dtype.offset_type is not None @@ -1320,7 +1320,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: offset_provider_t = self.subgraph_builder.get_offset_provider_type( self.subgraph_builder.connectivity_key_over(offset_type) ) - assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_t, gtx_common.NeighborTableType) input_conn_types[offset_type] = offset_provider_t if len(input_conn_types) == 0: @@ -1379,7 +1379,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: conn_desc = self.sdfg.arrays[conn_data] conn_desc.transient = False - origin_map_index = gtir_to_sdfg_utils.get_map_variable(conn_type.source_dim) + origin_map_index = gtir_to_sdfg_utils.get_map_variable(conn_type.domain[0]) # The layout of connectivity tables is known. assert len(conn_type.domain) == 2 @@ -1440,7 +1440,7 @@ def _broadcast_const_list( offset_provider_t = self.subgraph_builder.get_offset_provider_type( self.subgraph_builder.connectivity_key_over(list_type.offset_type) ) - assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_t, gtx_common.NeighborTableType) local_size = offset_provider_t.max_neighbors map_index = gtir_to_sdfg_utils.get_map_variable(list_type.offset_type) @@ -1481,7 +1481,7 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: offset_provider_type = self.subgraph_builder.get_offset_provider_type( self.subgraph_builder.connectivity_key_over(offset_type) ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) inp_conn = "_in" outp_conn = "_out" @@ -1509,7 +1509,7 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: ) self.state.add_node(reduce_node) - origin_map_index = gtir_to_sdfg_utils.get_map_variable(offset_provider_type.source_dim) + origin_map_index = gtir_to_sdfg_utils.get_map_variable(offset_provider_type.domain[0]) self._add_input_data_edge( self.state.add_access(connectivity), dace_subsets.Range.from_string( @@ -1699,7 +1699,7 @@ def _make_dynamic_neighbor_offset( def _make_unstructured_shift( self, it: IteratorExpr, - conn_type: gtx_common.NeighborConnectivityType, + conn_type: gtx_common.NeighborTableType, conn_node: dace_nodes.AccessNode, offset_expr: DataExpr, ) -> IteratorExpr: @@ -1707,19 +1707,19 @@ def _make_unstructured_shift( # make sure that the field can be dereferenced with the given connectivity type assert any(dim == conn_type.codomain for dim, _ in it.field_domain) # make sure that the iterator can access the connectivity table - assert conn_type.source_dim in it.indices - conn_source_index = it.indices[conn_type.source_dim] + assert conn_type.domain[0] in it.indices + conn_source_index = it.indices[conn_type.domain[0]] assert isinstance(conn_source_index, SymbolExpr) shifted_indices = { - dim: idx for dim, idx in it.indices.items() if dim != conn_type.source_dim + dim: idx for dim, idx in it.indices.items() if dim != conn_type.domain[0] } if isinstance(offset_expr, SymbolExpr): # use memlet to retrieve the neighbor index shifted_indices[conn_type.codomain] = MemletExpr( dc_node=conn_node, gt_field=ts.FieldType( - dims=[conn_type.source_dim], + dims=[conn_type.domain[0]], dtype=ts.ListType( element_type=tt.from_dtype(conn_type.dtype), offset_type=gtx_common.ConstList, @@ -1761,7 +1761,7 @@ def _visit_shift(self, node: gtir.FunCall) -> IteratorExpr: offset_provider_type = self.subgraph_builder.get_offset_provider_type( offset_provider_arg.value ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) # a named offset → unstructured shift; the offset value may be a static # `OffsetLiteral` or a dynamic offset (handled by `_make_unstructured_shift`). # initially, the storage for the connectivity tables is created as transient; diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index 50a445146c..0079e47638 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -328,7 +328,7 @@ def _construct_if_branch_output( offset_provider_type = sdfg_builder.get_offset_provider_type( sdfg_builder.connectivity_key_over(out_type.dtype.offset_type) ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) shape = [*shape, offset_provider_type.max_neighbors] out, _ = sdfg_builder.add_temp_array(ctx.sdfg, shape, dtype) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index b158097a51..911078c9c7 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -387,7 +387,7 @@ def get_scan_output_shape( offset_provider_type = sdfg_builder.get_offset_provider_type( sdfg_builder.connectivity_key_over(offset_type) ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) list_size = offset_provider_type.max_neighbors return [scan_column_size, dace.symbolic.SymExpr(list_size)] diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py index c614719aff..55d3c77ada 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -60,7 +60,7 @@ def connectivity_identifier(name: str) -> str: def is_connectivity_identifier( - name: str, offset_provider_type: gtx_common.OffsetProviderType | None = None + name: str, offset_provider_type: gtx_common.TableTypes | None = None ) -> bool: if (m := CONNECTIVITY_INDENTIFIER_RE.match(name)) is None: return False @@ -76,7 +76,7 @@ def _field_symbol( field_name: str, dim: gtx_common.Dimension, sym: Literal["size", "stride"], - offset_provider_type: gtx_common.OffsetProviderType | None, + offset_provider_type: gtx_common.TableTypes | None, ) -> dace.symbol: if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is None: name = f"__{field_name}_{gtx_common.codegen_name(dim.tag)}_{sym}" @@ -85,10 +85,10 @@ def _field_symbol( offset = gtx_common.from_codegen_name(m[1]) assert offset in offset_provider_type conn_type = offset_provider_type[offset] - assert isinstance(conn_type, gtx_common.NeighborConnectivityType) - if dim == conn_type.source_dim: + assert isinstance(conn_type, gtx_common.NeighborTableType) + if dim == conn_type.domain[0]: name = f"__{field_name}_source_{sym}" - elif dim == conn_type.neighbor_dim: + elif dim == conn_type.domain[1]: name = f"__{field_name}_neighbor_{sym}" else: raise ValueError(f"Unexpect dimension '{dim}' for '{offset}' connectivity.") @@ -98,7 +98,7 @@ def _field_symbol( def field_size_symbol( field_name: str, dim: gtx_common.Dimension, - offset_provider_type: gtx_common.OffsetProviderType, + offset_provider_type: gtx_common.TableTypes, ) -> dace.symbol: return _field_symbol(field_name, dim, "size", offset_provider_type) @@ -106,7 +106,7 @@ def field_size_symbol( def field_stride_symbol( field_name: str, dim: gtx_common.Dimension, - offset_provider_type: gtx_common.OffsetProviderType | None = None, + offset_provider_type: gtx_common.TableTypes | None = None, ) -> dace.symbol: return _field_symbol(field_name, dim, "stride", offset_provider_type) @@ -114,7 +114,7 @@ def field_stride_symbol( def local_dimension_size( field_name: str, dim: gtx_common.Dimension, - neighbor_table_types: dict[str, gtx_common.NeighborConnectivityType], + neighbor_table_types: dict[str, gtx_common.NeighborTableType], ) -> int: """ Number of neighbors along the local dimension `dim` of the field or connectivity table. @@ -125,7 +125,7 @@ def local_dimension_size( """ if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is not None: own_type = neighbor_table_types[gtx_common.from_codegen_name(m[1])] - if own_type.neighbor_dim == dim: + if own_type.domain[1] == dim: return own_type.max_neighbors return neighbor_table_types[ gtx_common.connectivity_key_over(neighbor_table_types, dim) @@ -151,15 +151,15 @@ def range_stop_symbol(field_name: str, dim: gtx_common.Dimension) -> dace.symbol def filter_connectivity_types( - offset_provider_type: gtx_common.OffsetProviderType, -) -> dict[str, gtx_common.NeighborConnectivityType]: + offset_provider_type: gtx_common.TableTypes, +) -> dict[str, gtx_common.NeighborTableType]: """ - Filter offset provider types of type `NeighborConnectivityType`. + Filter offset provider types of type `NeighborTableType`. In other words, filter out the cartesian offset providers. """ return { offset: conn for offset, conn in offset_provider_type.items() - if isinstance(conn, gtx_common.NeighborConnectivityType) + if isinstance(conn, gtx_common.NeighborTableType) } diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index ee96fbeff2..df2ceb6670 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -31,7 +31,7 @@ def find_constant_symbols( ir: itir.Program, sdfg: dace.SDFG, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, disable_field_origin_on_program_arguments: bool, unstructured_horizontal_has_unit_stride: bool, ) -> dict[str, int]: @@ -56,13 +56,13 @@ def find_constant_symbols( # Same for connectivity tables, for which the first dimension is always horizontal for offset, conn_type in offset_provider_type.items(): if ( - isinstance(conn_type, common.NeighborConnectivityType) + isinstance(conn_type, common.NeighborTableType) and (conn_id := gtx_dace_args.connectivity_identifier(offset)) in sdfg.arrays ): assert not sdfg.arrays[conn_id].transient - assert conn_type.source_dim.kind == common.DimensionKind.HORIZONTAL + assert conn_type.domain[0].kind == common.DimensionKind.HORIZONTAL sdfg_stride_symbol = gtx_dace_args.field_stride_symbol( - conn_id, conn_type.source_dim, offset_provider_type + conn_id, conn_type.domain[0], offset_provider_type ) constant_symbols[sdfg_stride_symbol.name] = 1 diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 453d5f8dad..6db624755c 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -533,7 +533,7 @@ def is_concretizable(symbol_type: ts.TypeSpec, to_type: ts.TypeSpec) -> bool: True >>> is_concretizable( - ... ts.DeferredType(constraint=ts.OffsetType), + ... ts.DeferredType(constraint=ts.ShiftType), ... to_type=ts.FieldType(dtype=ts.ScalarType(kind=ts.ScalarKind.BOOL), dims=[]), ... ) False @@ -669,19 +669,19 @@ def return_type_field( except ValueError as ex: raise ValueError("Could not deduce return type of invalid remap operation.") from ex - if not isinstance(with_args[0], ts.OffsetType): - raise ValueError(f"First argument must be of type '{ts.OffsetType}', got '{with_args[0]}'.") + if not isinstance(with_args[0], ts.ShiftType): + raise ValueError(f"First argument must be of type '{ts.ShiftType}', got '{with_args[0]}'.") - source_dim = with_args[0].source - target_dims = with_args[0].target + codomain = with_args[0].codomain + domain_dims = with_args[0].domain new_dims = [] # TODO: This code does not handle ellipses for dimensions. Fix it. assert field_type.dims is not ... for d in field_type.dims: - if d != source_dim: + if d != codomain: new_dims.append(d) else: - new_dims.extend(target_dims) + new_dims.extend(domain_dims) return ts.FieldType(dims=new_dims, dtype=field_type.dtype) @@ -890,10 +890,10 @@ def function_signature_incompatibilities_field( yield f"Function takes at least 1 argument, but {len(args)} were given." return for arg in args: - if not isinstance(arg, ts.OffsetType): - yield f"Expected arguments to be of type '{ts.OffsetType}', got '{arg}'." + if not isinstance(arg, ts.ShiftType): + yield f"Expected arguments to be of type '{ts.ShiftType}', got '{arg}'." return - if len(args) > 1 and len(arg.target) > 1: + if len(args) > 1 and len(arg.domain) > 1: yield f"Function takes only 1 argument in unstructured case, but {len(args)} were given." return @@ -901,15 +901,15 @@ def function_signature_incompatibilities_field( yield f"Got unexpected keyword argument(s) '{', '.join(kwargs.keys())}'." return - source_dim = args[0].source # type: ignore[attr-defined] # ensured by loop above - target_dims = args[0].target # type: ignore[attr-defined] # ensured by loop above + codomain = args[0].codomain # type: ignore[attr-defined] # ensured by loop above + domain_dims = args[0].domain # type: ignore[attr-defined] # ensured by loop above assert field_type.dims is not ... - if field_type.dims and source_dim not in field_type.dims: + if field_type.dims and codomain not in field_type.dims: yield ( f"Incompatible offset can not shift field defined on " f"{', '.join([dim.__qualname__ for dim in field_type.dims])} from " - f"{source_dim.__qualname__} to target dim(s): " - f"{', '.join([dim.tag for dim in target_dims])}" + f"{codomain.__qualname__} to target dim(s): " + f"{', '.join([dim.tag for dim in domain_dims])}" ) diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 5ad8ea7105..648dfa866b 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -67,16 +67,28 @@ def __str__(self) -> str: return f"Index[{self.dim}]" -class OffsetType(TypeSpec): - # TODO(havogt): replace by ConnectivityType - source: common.Dimension - target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] +class ShiftType(TypeSpec): + """ + The type of a shift: it takes a field over `codomain` to a field over `domain`. + + `domain` has one dimension for a Cartesian shift (`KDim + 1`) and for a single neighbor + (`V2E[i]`), and two -- the connectivity's domain and its local dimension -- for all + neighbors (`V2E`). + """ + + codomain: common.Dimension + domain: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] #: The offset-provider key; `None` for the untagged Cartesian `Dim + offset`. tag: Optional[common.Tag] = None def __str__(self) -> str: tag = "" if self.tag is None else f"{self.tag}: " - return f"Offset[{tag}{self.source}, {self.target}]" + domain = ( + str(self.domain[0]) + if len(self.domain) == 1 + else f"({', '.join(str(dim) for dim in self.domain)})" + ) + return f"Shift[{tag}{self.codomain} -> {domain}]" class ScalarKind(eve_types.IntEnum): diff --git a/src/gt4py/next/type_system/type_translation.py b/src/gt4py/next/type_system/type_translation.py index 9722b27f5b..d66214c992 100644 --- a/src/gt4py/next/type_system/type_translation.py +++ b/src/gt4py/next/type_system/type_translation.py @@ -365,7 +365,7 @@ def from_value(value: Any) -> ts.TypeSpec: type_ = xtyping.infer_type(value, annotate_callable_kwargs=True) symbol_type = from_type_hint(type_) - if isinstance(symbol_type, (ts.DataType, ts.CallableType, ts.OffsetType, ts.DimensionType)): + if isinstance(symbol_type, (ts.DataType, ts.CallableType, ts.ShiftType, ts.DimensionType)): return symbol_type else: raise ValueError(f"Impossible to map '{value}' value to a 'Symbol'.") diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 4b73f163a9..efe47928ae 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -206,7 +206,7 @@ def sizes(self) -> tuple[int, int, int]: ... def offset_provider(self) -> common.OffsetProvider: ... @property - def offset_provider_type(self) -> common.OffsetProviderType: ... + def offset_provider_type(self) -> common.TableTypes: ... def simple_cartesian_grid( @@ -248,7 +248,7 @@ def num_edges(self) -> int: ... def offset_provider(self) -> common.OffsetProvider: ... @property - def offset_provider_type(self) -> common.OffsetProviderType: ... + def offset_provider_type(self) -> common.TableTypes: ... def simple_mesh(allocator) -> MeshDescriptor: diff --git a/tests/next_tests/integration_tests/feature_tests/instrumentation_tests/test_hooks.py b/tests/next_tests/integration_tests/feature_tests/instrumentation_tests/test_hooks.py index ce16667142..5e750f3c65 100644 --- a/tests/next_tests/integration_tests/feature_tests/instrumentation_tests/test_hooks.py +++ b/tests/next_tests/integration_tests/feature_tests/instrumentation_tests/test_hooks.py @@ -214,7 +214,7 @@ def custom_compile_variant_hook( key: gtx_typing.CompiledProgramsKey, backend: gtx_typing.Backend, argument_descriptors: dict[type, dict[str, Any]], - offset_provider: common.OffsetProviderType | common.OffsetProvider, + offset_provider: common.TableTypes | common.OffsetProvider, ) -> None: callback_results.append( ( diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py index bd675b5f51..d5406e4e08 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py @@ -10,6 +10,7 @@ import pytest import gt4py.next as gtx +from gt4py.next import common from gt4py.next.iterator.builtins import deref, named_range, shift, unstructured_domain, as_fieldop from gt4py.next.iterator.runtime import set_at, fendef, fundef, offset @@ -53,7 +54,10 @@ def test_strided_offset_provider(program_processor): program_processor, validate = program_processor LocA_size = 2 - max_neighbors = LocA2LocAB_offset_provider.__gt_type__().max_neighbors + # the table's type as bound under its key: a table's own `__gt_type__()` is only structural + max_neighbors = common.offset_provider_to_type({"O": LocA2LocAB_offset_provider})[ + "O" + ].max_neighbors LocAB_size = LocA_size * max_neighbors rng = np.random.default_rng() diff --git a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py index 82ba0ecbd4..279a5ef023 100644 --- a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py +++ b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py @@ -137,7 +137,7 @@ def test_shift_tag_differs_from_local_dim_name(case_tag_vs_local_dim): """ Ensure a shift works with an offset tag that differs from the local dimension's name. - If `NeighborConnectivityType.neighbor_dim` did not match the `FieldOffset` value, + If the local dimension of the `NeighborTableType` did not match the `FieldOffset` value, gtfn would silently ignore the neighbor index, see https://github.com/GridTools/gridtools/pull/1814. """ diff --git a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py index 3c7157166d..1531a2b09c 100644 --- a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py +++ b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py @@ -315,7 +315,7 @@ def test_unstructured_translate(shift_chain, expected_end_domain): def test_unstructured_translate_with_symbolic_domain_sizes(as_type): # With `symbolic_domain_sizes` the translated range is taken from the provided size # expression instead of the connectivity table. This makes `translate` work for a type-only - # `OffsetProviderType` (which has no table) as well as a runtime `OffsetProvider`. + # `TableTypes` (which has no table) as well as a runtime `OffsetProvider`. offset_provider = { V2EDim.tag: constructors.as_connectivity( domain={Vertex: (0, 4), V2EDim: 1}, diff --git a/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py b/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py index a742561830..d933548042 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py @@ -29,11 +29,16 @@ def foo(inp): return deref(inp) -connectivity = common.ConnectivityType( - domain=[dummy_origin, dummy_neighbor], - codomain=dummy_codomain, +connectivity = common.NeighborTableType( + connectivity=common.ConnectivityType( + domain=(dummy_origin, dummy_neighbor), + codomain=dummy_codomain, + skip_value=common._DEFAULT_SKIP_VALUE, + dtype=None, + ), skip_value=common._DEFAULT_SKIP_VALUE, dtype=None, + max_neighbors=1, ) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py index f104ed3b40..733b4f1434 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py @@ -34,10 +34,15 @@ class Dim2(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... def dummy_connectivity_type(max_neighbors: int, has_skip_values: bool): - return common.NeighborConnectivityType( - domain=[dummy_origin, dummy_neighbor], - codomain=dummy_codomain, - skip_value=common._DEFAULT_SKIP_VALUE if has_skip_values else None, + skip_value = common._DEFAULT_SKIP_VALUE if has_skip_values else None + return common.NeighborTableType( + connectivity=common.ConnectivityType( + domain=(dummy_origin, dummy_neighbor), + codomain=dummy_codomain, + skip_value=skip_value, + dtype=None, + ), + skip_value=skip_value, dtype=None, max_neighbors=max_neighbors, ) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py index 50141db06b..96e121ac4f 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py @@ -29,12 +29,14 @@ def test_local_dimension_size(): from next_tests.toy_connectivity import V2EDim, Vertex, Edge - def conn_type(max_neighbors: int) -> common.NeighborConnectivityType: - return common.NeighborConnectivityType( - domain=(Vertex, V2EDim), - codomain=Edge, + def conn_type(max_neighbors: int) -> common.NeighborTableType: + dtype = core_defs.dtype(np.int32) + return common.NeighborTableType( + connectivity=common.ConnectivityType( + domain=(Vertex, V2EDim), codomain=Edge, skip_value=None, dtype=dtype + ), skip_value=None, - dtype=core_defs.dtype(np.int32), + dtype=dtype, max_neighbors=max_neighbors, ) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 3ff9bcc055..7f975c7d2f 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -20,7 +20,7 @@ DimensionKind, LocalDimensionIndex, NeighborConnectivity, - NeighborConnectivityType, + NeighborTableType, ) from gt4py.next.ffront import transform_utils from gt4py.next.type_system import type_specifications as ts, type_translation @@ -73,7 +73,7 @@ def _declare(source: str) -> dict: class TestDeclaration: def test_owner_and_dimensions(self): assert V2E.Local.owner is V2E - assert V2E.origin is Vertex + assert V2E.domain is Vertex assert V2E.codomain is Edge assert V2E.Local.kind is DimensionKind.LOCAL assert issubclass(V2E.Local, DimensionIndex) @@ -160,21 +160,21 @@ class C(NeighborConnectivity[Edge, Edge]): class C(NeighborConnectivity): class Local(LocalDimensionIndex): ... """, - "must derive from 'NeighborConnectivity\\[Origin, Codomain\\]'", + "must derive from 'NeighborConnectivity\\[Domain, Codomain\\]'", ), ( """ class C(V2E): class Local(LocalDimensionIndex): ... """, - "must derive from 'NeighborConnectivity\\[Origin, Codomain\\]'", + "must derive from 'NeighborConnectivity\\[Domain, Codomain\\]'", ), ( """ class C(NeighborConnectivity[V2E.Local, Edge]): class Local(LocalDimensionIndex): ... """, - "'Origin' must be a non-local dimension", + "'Domain' must be a non-local dimension", ), ( """ @@ -292,10 +292,13 @@ def _table_type( max_neighbors=4, skip_value=common._DEFAULT_SKIP_VALUE, dtype=np.int32, -) -> NeighborConnectivityType: - return NeighborConnectivityType( - domain=domain, - codomain=codomain, +) -> NeighborTableType: + """The type of a table no declaration names, as `NeighborTable.__gt_type__()` describes it.""" + structure = common.ConnectivityType( + domain=domain, codomain=codomain, skip_value=skip_value, dtype=core_defs.dtype(dtype) + ) + return NeighborTableType( + connectivity=structure, skip_value=skip_value, dtype=core_defs.dtype(dtype), max_neighbors=max_neighbors, @@ -373,13 +376,100 @@ class Local(LocalDimensionIndex): ... ) +class TestNeighborTableType: + @staticmethod + def _v2e_shaped_table(codomain=Edge): + from gt4py.next import constructors + + return constructors.as_connectivity( + domain={Vertex: 2, V2E.Local: 4}, + codomain=codomain, + data=np.array([[0, 1, 2, -1], [1, 2, 3, 0]]), + skip_value=common._DEFAULT_SKIP_VALUE, + ) + + def test_domain_and_codomain_come_from_the_declaration(self): + table_type = common.check_neighbor_table(V2E, self._v2e_shaped_table()) + assert table_type.connectivity is V2E + assert table_type.domain == (Vertex, V2E.Local) + assert table_type.codomain is Edge + assert table_type.max_neighbors == 4 and table_type.has_skip_values + + def test_a_table_alone_has_its_structural_type(self): + structure = self._v2e_shaped_table().__gt_type__() + assert type(structure) is common.ConnectivityType + assert structure.domain == (Vertex, V2E.Local) and structure.codomain is Edge + + def test_bound_by_the_provider_key(self): + # a sharer's table looks like its owner's but for the codomain: only the key tells + # which declaration it is bound to + sharer = _declare( + """ + class V2V(NeighborConnectivity[Vertex, Vertex]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2V"] + v2e_table, v2v_table = self._v2e_shaped_table(), self._v2e_shaped_table(Vertex) + table_types = common.offset_provider_to_type( + {V2E.offset_tag: v2e_table, sharer.offset_tag: v2v_table, "undeclared": v2e_table} + ) + assert table_types[V2E.offset_tag].connectivity is V2E + assert table_types[sharer.offset_tag].connectivity is sharer + assert table_types[sharer.offset_tag].domain == table_types[V2E.offset_tag].domain + assert table_types[sharer.offset_tag].codomain is Vertex + # a key no declaration answers to leaves the table typed by its structure + assert table_types["undeclared"].connectivity == v2e_table.__gt_type__() + assert table_types["undeclared"].domain == table_types[V2E.offset_tag].domain + + def test_a_table_that_does_not_match_its_key(self): + with pytest.raises(ValueError, match="its codomain is"): + common.offset_provider_to_type({V2E.offset_tag: self._v2e_shaped_table(Vertex)}) + + def test_a_type_bound_to_another_declaration(self): + sharer = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2EShared"] + v2e_type = common.check_neighbor_table(V2E, self._v2e_shaped_table()) + with pytest.raises(ValueError, match="bound to 'V2E'"): + common.check_neighbor_table(sharer, v2e_type) + + def test_fingerprint_tells_the_declarations_apart(self): + from gt4py.next import fingerprinting + + sharer = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2EShared"] + table = self._v2e_shaped_table() + owner_type = common.check_neighbor_table(V2E, table) + sharer_type = common.check_neighbor_table(sharer, table) + # lenient: `_declare` classes are not importable + assert fingerprinting.lenient_fingerprinter( + owner_type + ) != fingerprinting.lenient_fingerprinter(sharer_type) + assert fingerprinting.lenient_fingerprinter( + owner_type + ) == fingerprinting.lenient_fingerprinter(common.check_neighbor_table(V2E, table)) + + class TestFrontendIntegration: def test_from_value_is_an_offset(self): # NOTE: pins the `__gt_type__` branch of `from_value` ahead of the dimension branch; a # connectivity declaration is a class, like a dimension. - assert type_translation.from_value(V2E) == ts.OffsetType( - source=Edge, target=(Vertex, V2E.Local), tag=V2E.Local.tag + assert type_translation.from_value(V2E) == ts.ShiftType( + codomain=Edge, domain=(Vertex, V2E.Local), tag=V2E.Local.tag + ) + + def test_shift_type_str(self): + assert str(V2E.__gt_type__()) == ( + f"Shift[{V2E.Local.tag}: {Edge} -> ({Vertex}, {V2E.Local})]" ) + assert str(ts.ShiftType(codomain=KDim, domain=(KDim,))) == f"Shift[{KDim} -> {KDim}]" def test_field_offset_is_derived_once(self): assert V2E.__gt_field_offset__() is V2E.__gt_field_offset__() @@ -406,11 +496,13 @@ def test_attribute_errors_are_dsl_errors(self): from gt4py.next.ffront.func_to_foast import FieldOperatorParser from gt4py.next import Dims, Field - def origin_of(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: - return a(V2E.origin) + def domain_of(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return a(V2E.domain) - with pytest.raises(errors.DSLError, match="has no attribute 'origin'"): - FieldOperatorParser.apply_to_function(origin_of) + # NOTE: `V2E` is typed as a `ts.ShiftType`, whose `domain` is a field of the type, not a + # value in DSL code + with pytest.raises(errors.DSLError, match="has no attribute 'domain'"): + FieldOperatorParser.apply_to_function(domain_of) def test_fingerprint_covers_the_declaration(self): from gt4py.next import fingerprinting @@ -493,7 +585,7 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): class TestConnectivityKeyOver: def _type(self, connectivity): - return _table_type(domain=(connectivity.origin, common.local_dimension_of(connectivity))) + return _table_type(domain=(connectivity.domain, common.local_dimension_of(connectivity))) def test_owner_is_preferred(self): ns = _declare( From 132df396d7f1b1838fa51bece757c616279bae3b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 00:03:26 +0200 Subject: [PATCH 10/13] docs[next]: renumber the connectivities-as-types ADR to 0030 ADR 0028 on main is now 'Plain Builders Instead of factory-boy Factories' (#2808), so the two ADRs of this stack move up by one. --- docs/development/ADRs/next/0019-Connectivities.md | 4 ++-- ...ectivities_As_Types.md => 0030-Connectivities_As_Types.md} | 0 docs/development/ADRs/next/README.md | 2 +- typing_tests/pyright_probes.py | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) rename docs/development/ADRs/next/{0029-Connectivities_As_Types.md => 0030-Connectivities_As_Types.md} (100%) diff --git a/docs/development/ADRs/next/0019-Connectivities.md b/docs/development/ADRs/next/0019-Connectivities.md index f22a13e039..0c80a6cf55 100644 --- a/docs/development/ADRs/next/0019-Connectivities.md +++ b/docs/development/ADRs/next/0019-Connectivities.md @@ -31,7 +31,7 @@ We update and introduce the following concepts **NeighborTable** is a _GatherConnectivity_ that is a 2D mapping of the N neighbors of a Location A to a Location B, backed by a buffer. -**ConnectivityType**, **NeighborTableType** contain all information that is needed for compilation. A `NeighborTableType` is the type of a table bound to a `NeighborConnectivity` declaration (ADR 0029). +**ConnectivityType**, **NeighborTableType** contain all information that is needed for compilation. A `NeighborTableType` is the type of a table bound to a `NeighborConnectivity` declaration (ADR 0030). ### Full definitions @@ -63,4 +63,4 @@ The only supported `Connectivity`s in compiled backends (currently) are `Neighbo ### 2026-09-24 -- `NeighborConnectivityType` is renamed `NeighborTableType` and typed by the `NeighborConnectivity` declaration its table is bound to; a `NeighborTable`'s own `__gt_type__()` is the structural `ConnectivityType` (ADR 0029). +- `NeighborConnectivityType` is renamed `NeighborTableType` and typed by the `NeighborConnectivity` declaration its table is bound to; a `NeighborTable`'s own `__gt_type__()` is the structural `ConnectivityType` (ADR 0030). diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0030-Connectivities_As_Types.md similarity index 100% rename from docs/development/ADRs/next/0029-Connectivities_As_Types.md rename to docs/development/ADRs/next/0030-Connectivities_As_Types.md diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index 4946d31972..db55239983 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -23,7 +23,7 @@ Writing a new ADR is simple: - [0023 - Fingerprinting](0023-Fingerprinting.md) - [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md) - [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md) -- [0029 - Connectivities as Types](0029-Connectivities_As_Types.md) +- [0030 - Connectivities as Types](0030-Connectivities_As_Types.md) ### Frontend and Parsing #frontend diff --git a/typing_tests/pyright_probes.py b/typing_tests/pyright_probes.py index e9be2423d9..fbc47762a0 100644 --- a/typing_tests/pyright_probes.py +++ b/typing_tests/pyright_probes.py @@ -12,7 +12,7 @@ The cases in `test_next.yaml` run under mypy only, and the two checkers disagree about what counts as a type: an annotated `Local` on a connectivity or its metaclass makes every declaration's local dimension a *variable* for pyright, so `Field[Dims[V, V2E.Local], float]` -is rejected there while mypy accepts it (see ADR 0029). Everything here must be error-free. +is rejected there while mypy accepts it (see ADR 0030). Everything here must be error-free. """ from __future__ import annotations From 90547fb073f78d1ad4a50da45790e91c7bc19089 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 00:04:25 +0200 Subject: [PATCH 11/13] docs[next]: refer to the dimensions-as-nominal-types ADR as 0029 --- docs/development/ADRs/next/0030-Connectivities_As_Types.md | 2 +- docs/development/ADRs/next/README.md | 2 +- .../runners/dace/lowering/gtir_to_sdfg_lambda.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/development/ADRs/next/0030-Connectivities_As_Types.md b/docs/development/ADRs/next/0030-Connectivities_As_Types.md index a08b1555f7..60f01393e0 100644 --- a/docs/development/ADRs/next/0030-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0030-Connectivities_As_Types.md @@ -23,7 +23,7 @@ def f(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: ``` The declaration is written in DSL code, owns its local dimension, and states the -constraints a neighbor table bound to it has to satisfy. It holds no data. It builds on [ADR 0028](0028-Dimensions_As_Nominal_Types.md): the +constraints a neighbor table bound to it has to satisfy. It holds no data. It builds on [ADR 0029](0029-Dimensions_As_Nominal_Types.md): the connectivity, like a dimension, is identified by its type, and `V2E.Local` is an ordinary dimension class with the tag `.V2E.Local`. diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index db55239983..35c94b5b78 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -22,7 +22,7 @@ Writing a new ADR is simple: - [0021 - Argument Descriptors](0021-Argument-Descriptors.md) - [0023 - Fingerprinting](0023-Fingerprinting.md) - [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md) -- [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md) +- [0029 - Dimensions as Nominal Types](0029-Dimensions_As_Nominal_Types.md) - [0030 - Connectivities as Types](0030-Connectivities_As_Types.md) ### Frontend and Parsing #frontend diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 98e5d7a355..7cbe8ae87d 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -1154,7 +1154,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: # NOTE: the connectivity's own local dimension, not one synthesized from the offset # tag. The latter named a local dimension after the *offset*, which only coincided with # the real one under the old `V2EDim = Dimension("V2E")` convention, and under nominal - # identity (ADR 0028) a tag string cannot be turned back into a dimension at all. + # identity (ADR 0029) a tag string cannot be turned back into a dimension at all. offset_type = conn_type.domain[1] neighbor_idx = gtir_to_sdfg_utils.get_map_variable(offset_type) From 7ea88862000195bdcdc81678a2cc1121640e8efb Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 00:29:38 +0200 Subject: [PATCH 12/13] refactor[next]: localness is the class; remove DimensionKind.LOCAL MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A dimension is local if and only if it subclasses `LocalDimensionIndex` (`common.is_local_dimension`), so `DimensionKind` loses its `LOCAL` member and is `HORIZONTAL | VERTICAL`. A local dimension's `kind` is `None`, and declaring one with `kind=` is a `TypeError`; `None` rather than `HORIZONTAL` keeps every backend comparison against `HORIZONTAL` / `VERTICAL` meaning what it meant. `order_dimensions` sorts by an explicit rank (horizontal, local, vertical), so sparse-field layouts do not move. Displays derive the label from the class: `Local[local]`, the `ₗ` IR suffix and DaCe's `_gtx_localdim` map variables are unchanged. The tree's `DimensionIndex(kind=LOCAL)` declarations become owner-less `LocalDimensionIndex` subclasses. pyright probes pin the Cartesian axis levels of ADR 0029, including staggering and shifting a local dimension, with `reportUnnecessaryTypeIgnoreComment` so a rejection that stops firing fails the session. ADR 0030 records the change. --- .../ADRs/next/0030-Connectivities_As_Types.md | 28 ++++++- docs/user/next/QuickstartGuide.md | 4 +- docs/user/next/workshop/exercises/helpers.py | 19 +++-- docs/user/next/workshop/slides/slides_2.ipynb | 2 +- src/gt4py/next/common.py | 75 +++++++++++++------ src/gt4py/next/constructors.py | 2 +- src/gt4py/next/custom_layout_allocators.py | 2 +- src/gt4py/next/embedded/nd_array_field.py | 4 +- src/gt4py/next/ffront/fbuiltins.py | 4 +- .../ffront/foast_passes/type_deduction.py | 4 +- src/gt4py/next/ffront/past_to_itir.py | 2 +- src/gt4py/next/ffront/transform_utils.py | 2 +- src/gt4py/next/iterator/embedded.py | 2 +- src/gt4py/next/iterator/ir.py | 5 +- src/gt4py/next/iterator/pretty_printer.py | 10 ++- .../iterator/type_system/type_synthesizer.py | 4 +- .../codegens/gtfn/gtfn_module.py | 2 +- .../runners/dace/lowering/gtir_to_sdfg.py | 6 +- .../dace/lowering/gtir_to_sdfg_lambda.py | 12 +-- .../dace/lowering/gtir_to_sdfg_primitives.py | 2 +- .../dace/lowering/gtir_to_sdfg_types.py | 2 +- .../dace/lowering/gtir_to_sdfg_utils.py | 5 +- src/gt4py/next/type_system/type_info.py | 6 +- .../integration_tests/cases_utils.py | 2 +- ..._write_back_buffer_elimination_lowering.py | 2 +- .../iterator_tests/test_builtins.py | 2 +- .../test_strided_offset_provider.py | 2 +- .../test_offset_dimensions_names.py | 4 +- tests/next_tests/toy_connectivity.py | 8 +- .../embedded_tests/test_nd_array_field.py | 11 +-- .../test_decorator_domain_deduction.py | 2 +- .../ffront_tests/test_foast_to_gtir.py | 4 +- .../ffront_tests/test_source_utils.py | 4 +- .../ffront_tests/test_type_deduction.py | 5 +- .../ir_utils_tests/test_domain_utils.py | 6 +- .../test_embedded_field_with_list.py | 2 +- .../iterator_tests/test_pretty_printer.py | 2 +- .../transforms_tests/test_domain_inference.py | 2 +- .../transforms_tests/test_fuse_as_fieldop.py | 2 +- .../test_prune_empty_concat_where.py | 2 +- .../transforms_tests/test_unroll_reduce.py | 4 +- .../unit_tests/otf_tests/test_runners.py | 2 +- tests/next_tests/unit_tests/test_common.py | 13 ++-- .../test_custom_layout_allocators.py | 4 +- .../unit_tests/test_neighbor_connectivity.py | 46 ++++++++++-- .../type_system_tests/test_type_info.py | 5 +- typing_tests/pyright_probes.py | 35 ++++++++- typing_tests/pyrightconfig.json | 3 +- 48 files changed, 257 insertions(+), 121 deletions(-) diff --git a/docs/development/ADRs/next/0030-Connectivities_As_Types.md b/docs/development/ADRs/next/0030-Connectivities_As_Types.md index 60f01393e0..07ea6992b4 100644 --- a/docs/development/ADRs/next/0030-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0030-Connectivities_As_Types.md @@ -7,7 +7,7 @@ tags: [] - **Status**: proposed - **Authors**: Enrique González Paredes (@egparedes) - **Created**: 2026-09-21 -- **Updated**: 2026-09-24 +- **Updated**: 2026-10-02 A neighbor connectivity is declared as a **class**, and its local dimension as a class **nested** in it: @@ -128,9 +128,29 @@ accepted a `FieldOffset`, which is not a `Connectivity` either. A separate root would force every `type[DimensionIndex]` annotation in the tree (`ts.FieldType.dims`, `Domain`, `ConnectivityType.domain`, ...) to widen, and would then accept local dimensions wherever a primary one is meant anyway. The -tree already distinguishes local dimensions by a runtime `kind` check, so it -keeps doing so; generic constructors whose parameter must be a primary dimension -(`NeighborConnectivity[Domain, Codomain]`, `Staggered[D]`) check it at runtime. +tree tells local dimensions apart at runtime, and constructors whose parameter +must be a primary dimension (`NeighborConnectivity[Domain, Codomain]`) check it; +`Staggered[D]` rejects a local dimension statically too, through its bound on a +declared Cartesian axis (ADR 0029). + +### Localness is the class, and `DimensionKind.LOCAL` is removed + +A dimension is local if and only if it subclasses `LocalDimensionIndex` +(`common.is_local_dimension(dim)`), so `DimensionKind.LOCAL` is removed and +`DimensionKind` is `HORIZONTAL | VERTICAL`. A local dimension's `kind` is `None`, +and declaring one with `kind=` is a `TypeError`. `None` rather than `HORIZONTAL` +keeps every `kind == HORIZONTAL` / `kind != VERTICAL` comparison in the backends +meaning what it meant: with `HORIZONTAL`, a sparse field would silently count its +local axis as horizontal. Two consequences: + +- `None` does not order against the enum, so `order_dimensions` sorts by an + explicit rank — horizontal, then local, then vertical — the order `kind` used to + encode. It must not move, or the memory layout of sparse fields changes with it. +- Displays derive the label from the class: `str(V2E.Local)` is still + `Local[local]`, the IR pretty printer still marks a local axis with `ₗ`, and DaCe + map variables of a local dimension keep their `_gtx_localdim` suffix. + +What remains of `kind` is the layout sort key and the scan axis. ### `Local` is not annotated anywhere diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 63d5c2dda3..1695c0a63f 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -227,7 +227,7 @@ Another way to look at it is that transform uses the edge-to-cell connectivity t You can use the field offset `E2C` below to transform a field over cells to a field over edges using the edge-to-cell connectivities: ```{code-cell} ipython3 -class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2CDim(gtx.LocalDimensionIndex): ... E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim)) ``` @@ -379,7 +379,7 @@ print("where nested tuple return: {}".format(((result_1.asnumpy(), result_2.asnu As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to define the local dimension, the field offset, and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: ```{code-cell} ipython3 -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(gtx.LocalDimensionIndex): ... C2E = gtx.FieldOffset(C2EDim.tag, source=EdgeDim, target=(CellDim, C2EDim)) C2E_offset_provider = gtx.as_connectivity([CellDim, C2EDim], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index 82c0ef6a57..272fbdcc57 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -11,7 +11,14 @@ import gt4py.next as gtx from gt4py.next.iterator.embedded import MutableLocatedField from gt4py.next import neighbor_sum, where, Dims -from gt4py.next import CartesianAxisIndex, Dimension, DimensionIndex, DimensionKind, FieldOffset +from gt4py.next import ( + CartesianAxisIndex, + Dimension, + DimensionIndex, + LocalDimensionIndex, + DimensionKind, + FieldOffset, +) from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( run_gtfn as gtfn_cpu, @@ -389,31 +396,31 @@ class E(DimensionIndex): ... class K(CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... C2E = FieldOffset(C2EDim.tag, source=E, target=(C, C2EDim)) -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... V2E = FieldOffset(V2EDim.tag, source=E, target=(V, V2EDim)) -class E2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2VDim(LocalDimensionIndex): ... E2V = FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) -class E2CDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2CDim(LocalDimensionIndex): ... E2C = FieldOffset(E2CDim.tag, source=C, target=(E, E2CDim)) -class E2C2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2VDim(LocalDimensionIndex): ... E2C2V = FieldOffset(E2C2VDim.tag, source=V, target=(E, E2C2VDim)) diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index 6a44a6eef7..a336ed664a 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -273,7 +273,7 @@ "metadata": {}, "outputs": [], "source": [ - "class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...\n", + "class E2CDim(gtx.LocalDimensionIndex): ...\n", "\n", "\n", "E2C = gtx.FieldOffset(E2CDim.tag, source=Cell, target=(Edge, E2CDim))" diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 35ae15d33d..471b4a7abb 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -40,6 +40,8 @@ TypeVarTuple, Unpack, cast, + get_args, + get_origin, overload, ) @@ -122,15 +124,36 @@ def from_codegen_name(name: str) -> Tag: @enum.unique class DimensionKind(StrEnum): + """ + The role of a non-local dimension: a field's layout sort key, and the scan axis. + + There is no `LOCAL` member: whether a dimension is local is a fact about its class + (`is_local_dimension`), and a local dimension's `kind` is `None` (ADR 0030). + """ + HORIZONTAL = "horizontal" VERTICAL = "vertical" - LOCAL = "local" def __str__(self) -> str: return self.value -_DIM_KIND_ORDER = {DimensionKind.HORIZONTAL: 0, DimensionKind.LOCAL: 1, DimensionKind.VERTICAL: 2} +def is_local_dimension(dim: Any) -> bool: + """Return whether `dim` is a local dimension, i.e. a subclass of `LocalDimensionIndex`.""" + return isinstance(dim, DimensionMeta) and issubclass(dim, LocalDimensionIndex) + + +def _dimension_rank(dim: Dimension) -> int: + # NOTE: an explicit rank rather than a sort on `kind`: a local dimension's `kind` is `None`, + # which does not order against the enum. Horizontal, then local, then vertical -- the order + # `kind` used to encode; changing it would change the memory layout of sparse fields. + if is_local_dimension(dim): + return 1 + return 2 if dim.kind is DimensionKind.VERTICAL else 0 + + +def _kind_label(dim: DimensionMeta) -> str: + return "local" if is_local_dimension(dim) else str(dim.kind) class DimensionMeta(type): @@ -142,7 +165,8 @@ class DimensionMeta(type): metaclass, so this is the only place they can live. """ - kind: DimensionKind + #: `None` for a local dimension, whose localness is its class (see `is_local_dimension`). + kind: Optional[DimensionKind] # NOTE: mandatory, not redundant. Python sets `__hash__ = None` on any class body that # defines `__eq__` without it -- metaclasses included -- and `__eq__` below stays for the @@ -178,12 +202,12 @@ def value(cls) -> NoReturn: ) def __repr__(cls) -> str: - return f"{cls.tag}[{cls.kind}]" + return f"{cls.tag}[{_kind_label(cls)}]" def __str__(cls) -> str: # NOTE: the unqualified name, so diagnostics stay readable. `tag` is identity, not a # display name; `repr` carries the module and disambiguates when it matters. - return f"{cls.__qualname__}[{cls.kind}]" + return f"{cls.__qualname__}[{_kind_label(cls)}]" # NOTE: the self-type restricts index arithmetic to a Cartesian axis for the type checkers: # both bind it correctly at every call site (`C + 1` is an error for a mesh location `C`), @@ -277,7 +301,7 @@ class DimensionIndex(metaclass=DimensionMeta): False """ - kind: ClassVar[DimensionKind] = DimensionKind.HORIZONTAL + kind: ClassVar[Optional[DimensionKind]] = DimensionKind.HORIZONTAL __slots__ = ("value",) @@ -1512,7 +1536,7 @@ def is_neighbor_table(obj: Any) -> TypeGuard[NeighborTable]: return ( len(domain_dims) == 2 and domain_dims[0].kind is DimensionKind.HORIZONTAL - and domain_dims[1].kind is DimensionKind.LOCAL + and is_local_dimension(domain_dims[1]) ) @@ -1752,8 +1776,8 @@ class GridType(StrEnum): def order_dimensions(dims: Iterable[Dimension]) -> list[Dimension]: """Find the canonical ordering of the dimensions in `dims`.""" - if sum(1 for dim in dims if dim.kind == DimensionKind.LOCAL) > 1: - raise ValueError("There are more than one dimension with DimensionKind 'LOCAL'.") + if sum(1 for dim in dims if is_local_dimension(dim)) > 1: + raise ValueError("There is more than one local dimension.") # NOTE: `__qualname__`, not `tag`. The tag is qualified, so ordering by it would make a # field's canonical dimension order depend on *which module* each dimension is declared in -- # moving a declaration would silently reorder a field's dimensions. The unqualified name keeps @@ -1762,7 +1786,7 @@ def order_dimensions(dims: Iterable[Dimension]) -> list[Dimension]: return sorted( dims, key=lambda dim: ( - _DIM_KIND_ORDER[dim.kind], + _dimension_rank(dim), as_non_staggered(dim).__qualname__, as_non_staggered(dim).tag, ), @@ -1793,7 +1817,7 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: Find an ordering of multiple lists of dimensions. The resulting list contains all unique dimensions from the input lists, - sorted first by dims_kind_order, i.e., `Dimension.kind` (`HORIZONTAL` < `LOCAL` < `VERTICAL`) and then + sorted first by horizontal < local < vertical (see `order_dimensions`) and then lexicographically by `Dimension.tag`. Examples: @@ -1801,8 +1825,8 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: >>> class I(CartesianAxisIndex, kind=DimensionKind.HORIZONTAL): ... >>> class J(CartesianAxisIndex, kind=DimensionKind.HORIZONTAL): ... >>> class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... - >>> class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... - >>> class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... + >>> class E2V(LocalDimensionIndex): ... + >>> class E2C(LocalDimensionIndex): ... >>> promote_dims([J, K], [I, K]) == [I, J, K] True >>> promote_dims([K, J], [I, K]) @@ -1814,7 +1838,7 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: >>> promote_dims([I, E2C], [E2V, K]) Traceback (most recent call last): ... - ValueError: There are more than one dimension with DimensionKind 'LOCAL'. + ValueError: There is more than one local dimension. """ for dims in dims_list: @@ -1905,8 +1929,6 @@ def __getitem__(cls, base: Dimension) -> Dimension: ) if not isinstance(base, DimensionMeta): raise TypeError(f"'Staggered' expects a dimension, got '{base!r}'.") - if base.kind is DimensionKind.LOCAL: - raise TypeError(f"'{base.__qualname__}' is a local dimension and cannot be staggered.") if is_staggered(base): raise TypeError( f"'{base.__qualname__}' is already staggered; a dimension cannot be staggered twice." @@ -2045,7 +2067,7 @@ def connectivity_for_cartesian_shift(dim: Dimension, offset: int | float) -> Car return CartesianConnectivity(dim, int(integral_offset), codomain=dim) -class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): +class LocalDimensionIndex(DimensionIndex): """ A local dimension: the axis that runs over the neighbors of one element. @@ -2054,8 +2076,11 @@ class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): coefficients of a fixed-size stencil: >>> class LsqCoeff(LocalDimensionIndex, size=3): ... - >>> LsqCoeff.kind, LsqCoeff.owner, LsqCoeff.max_neighbors - (, None, 3) + >>> LsqCoeff.owner, LsqCoeff.max_neighbors, str(LsqCoeff) + (None, 3, 'LsqCoeff[local]') + + A local dimension has no `kind`: it is `None`, and `is_local_dimension` reads localness + from the class. Neighbor counts are optional. A declared count is a constraint the bound table has to satisfy (see `check_neighbor_table`); an undeclared one is taken from the table. @@ -2063,6 +2088,8 @@ class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): __slots__ = () + kind: ClassVar[Optional[DimensionKind]] = None + #: The connectivity this dimension is the local axis of, or `None` if it indexes no table. #: Set by `NeighborConnectivity` when the connectivity is declared. owner: ClassVar[Optional[type[NeighborConnectivity]]] = None @@ -2086,9 +2113,9 @@ def __init_subclass__( kind: Optional[DimensionKind] = None, **kwargs: Any, ) -> None: - if kind is not None and kind is not DimensionKind.LOCAL: + if kind is not None: raise TypeError( - f"'{cls.__qualname__}' is a local dimension and cannot have kind '{kind}'." + f"'{cls.__qualname__}' is a local dimension and has no kind; got '{kind}'." ) super().__init_subclass__(**kwargs) # NOTE: reset rather than inherited: a subclass of an owned local dimension is a @@ -2256,9 +2283,9 @@ def __init_subclass__( " the IR by its qualified name, which has to be importable." ) params = [ - xtyping.get_args(base) + get_args(base) for base in cls.__dict__.get("__orig_bases__", ()) - if xtyping.get_origin(base) is NeighborConnectivity + if get_origin(base) is NeighborConnectivity ] if len(params) != 1 or len(params[0]) != 2: raise TypeError( @@ -2267,7 +2294,7 @@ def __init_subclass__( ) domain, codomain = params[0] for role, dim in (("Domain", domain), ("Codomain", codomain)): - if not isinstance(dim, DimensionMeta) or dim.kind is DimensionKind.LOCAL: + if not isinstance(dim, DimensionMeta) or is_local_dimension(dim): raise TypeError(f"'{name}': '{role}' must be a non-local dimension, got '{dim}'.") local = cls.__dict__.get("Local") diff --git a/src/gt4py/next/constructors.py b/src/gt4py/next/constructors.py index 2bef28dce0..201b686764 100644 --- a/src/gt4py/next/constructors.py +++ b/src/gt4py/next/constructors.py @@ -653,7 +653,7 @@ def as_connectivity( >>> from gt4py import next as gtx >>> class Vertex(gtx.DimensionIndex): ... >>> class Edge(gtx.DimensionIndex): ... - >>> class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + >>> class V2EDim(gtx.LocalDimensionIndex): ... >>> data = np.array([[0, 1], [1, 2], [2, 0]]) >>> conn = gtx.as_connectivity([Vertex, V2EDim], Edge, data) >>> conn.ndarray diff --git a/src/gt4py/next/custom_layout_allocators.py b/src/gt4py/next/custom_layout_allocators.py index 5eee3833ac..453aee0bad 100644 --- a/src/gt4py/next/custom_layout_allocators.py +++ b/src/gt4py/next/custom_layout_allocators.py @@ -160,7 +160,7 @@ def pos_of_kind(kind: common.DimensionKind) -> list[int]: horizontals = pos_of_kind(common.DimensionKind.HORIZONTAL) verticals = pos_of_kind(common.DimensionKind.VERTICAL) - locals_ = pos_of_kind(common.DimensionKind.LOCAL) + locals_ = [i for i, dim in enumerate(dims) if common.is_local_dimension(dim)] layout_map = [0] * len(dims) for i, pos in enumerate(horizontals + verticals + locals_): diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index 8a1ee0a79c..26504c332d 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -967,11 +967,11 @@ def _builtin_op( ) -> NdArrayField[common.DimsT, core_defs.ScalarT]: xp = field.array_ns - if not axis.kind == common.DimensionKind.LOCAL: + if not common.is_local_dimension(axis): raise ValueError("Can only reduce local dimensions.") if axis not in field.domain.dims: raise ValueError(f"Field can not be reduced as it doesn't have dimension '{axis}'.") - if len([d for d in field.domain.dims if d.kind is common.DimensionKind.LOCAL]) > 1: + if len([d for d in field.domain.dims if common.is_local_dimension(d)]) > 1: raise NotImplementedError( "Reducing a field with more than one local dimension is not supported." ) diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 5c9ab763bd..dc8552425c 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -490,7 +490,7 @@ def _cache(self) -> dict: return {} def __post_init__(self) -> None: - if len(self.target) == 2 and self.target[1].kind != common.DimensionKind.LOCAL: + if len(self.target) == 2 and not common.is_local_dimension(self.target[1]): raise ValueError("Second dimension in offset must be a local dimension.") def __gt_type__(self) -> ts.ShiftType: @@ -547,5 +547,5 @@ def is_cartesian_offset(offset: FieldOffset | ts.ShiftType) -> bool: len(shift_type.domain) == 1 and shift_type.codomain == shift_type.domain[0] and shift_type.codomain.kind == shift_type.domain[0].kind - and shift_type.domain[0].kind != common.DimensionKind.LOCAL + and not common.is_local_dimension(shift_type.domain[0]) ) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index a486a0765b..1ef443c17c 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -469,7 +469,7 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri ) from ex new_type = types[index] case ts.ShiftType(codomain=codomain, domain=(domain, local), tag=tag): - if not local.kind == DimensionKind.LOCAL: + if not common.is_local_dimension(local): raise errors.DSLError( new_value.location, "Second dimension in offset must be a local dimension." ) @@ -824,7 +824,7 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call: hints=[f"Give the displacement, e.g. '{arg!s}[1]'."], ) elif isinstance(new_func.type, ts.DimensionType): - assert new_func.type.dim.kind == DimensionKind.LOCAL + assert common.is_local_dimension(new_func.type.dim) return foast.Call( func=new_func, args=new_args, diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index d82ed6e201..595ac45afd 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -406,7 +406,7 @@ def _construct_itir_domain_arg( dim_stop, ) - if dim.kind == common.DimensionKind.LOCAL: + if common.is_local_dimension(dim): raise ValueError(f"common.Dimension '{dim.__qualname__}' must not be local.") domain_args.append( itir.FunCall( diff --git a/src/gt4py/next/ffront/transform_utils.py b/src/gt4py/next/ffront/transform_utils.py index 4e24d83881..2b441796a6 100644 --- a/src/gt4py/next/ffront/transform_utils.py +++ b/src/gt4py/next/ffront/transform_utils.py @@ -66,7 +66,7 @@ def _deduce_grid_type( ): deduced_grid_type = common.GridType.UNSTRUCTURED break - if isinstance(o, common.DimensionMeta) and o.kind == common.DimensionKind.LOCAL: + if isinstance(o, common.DimensionMeta) and common.is_local_dimension(o): deduced_grid_type = common.GridType.UNSTRUCTURED break diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index c54feb6eb8..73c5e10102 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -913,7 +913,7 @@ def _get_sparse_dimensions(axes: Sequence[common.Dimension]) -> list[common.Dime return [ axis for axis in axes - if isinstance(axis, common.DimensionMeta) and axis.kind == common.DimensionKind.LOCAL + if isinstance(axis, common.DimensionMeta) and common.is_local_dimension(axis) ] diff --git a/src/gt4py/next/iterator/ir.py b/src/gt4py/next/iterator/ir.py index 717188c8a9..bfb0c666ed 100644 --- a/src/gt4py/next/iterator/ir.py +++ b/src/gt4py/next/iterator/ir.py @@ -99,9 +99,10 @@ def dim(self) -> common.Dimension: return common.resolve(self.value) @property - def kind(self) -> common.DimensionKind: + def kind(self) -> Optional[common.DimensionKind]: # NOTE: derived, not stored: the dimension class carries its kind, so a stored copy could - # only disagree with it (it used to, for local dimensions printed as vertical). + # only disagree with it (it used to, for local dimensions printed as vertical). `None` for a + # local dimension, see `common.is_local_dimension`. return self.dim.kind diff --git a/src/gt4py/next/iterator/pretty_printer.py b/src/gt4py/next/iterator/pretty_printer.py index a7376e0d22..201044177c 100644 --- a/src/gt4py/next/iterator/pretty_printer.py +++ b/src/gt4py/next/iterator/pretty_printer.py @@ -138,8 +138,8 @@ def implied_literal_type(value: str) -> ts.ScalarType: _AXIS_KIND_SUFFIX: Final = { common.DimensionKind.HORIZONTAL: "ₕ", common.DimensionKind.VERTICAL: "ᵥ", - common.DimensionKind.LOCAL: "ₗ", } +_LOCAL_AXIS_SUFFIX: Final = "ₗ" class PrettyPrinter(NodeTranslator): @@ -240,7 +240,13 @@ def visit_AxisLiteral(self, node: ir.AxisLiteral, *, prec: int) -> list[str]: dim: Optional[common.Dimension] = node.type.dim else: dim = common.resolve_loaded(node.value) - kind = _AXIS_KIND_SUFFIX[dim.kind] if dim is not None else "ₕ" + if dim is None: + kind = "ₕ" + elif common.is_local_dimension(dim): + kind = _LOCAL_AXIS_SUFFIX + else: + assert dim.kind is not None + kind = _AXIS_KIND_SUFFIX[dim.kind] return [str(node.value) + kind] def visit_SymRef(self, node: ir.SymRef, *, prec: int) -> list[str]: diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 859d8ae61d..910dd68dcb 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -406,7 +406,7 @@ def _canonicalize_nb_fields( Examples: >>> class Vertex(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> input_field = ts.FieldType( ... dims=[ ... Vertex, @@ -434,7 +434,7 @@ def _canonicalize_nb_fields( defined_dims = [] neighbor_dim = None for dim in input_dims: - if dim.kind == common.DimensionKind.LOCAL: + if common.is_local_dimension(dim): assert neighbor_dim is None neighbor_dim = dim else: diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index 45bec4de35..3c4f2ce1f2 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -86,7 +86,7 @@ def _process_regular_arguments( isinstance( dim, fbuiltins.FieldOffset ) # TODO(havogt): remove support for FieldOffset as Dimension - or dim.kind is common.DimensionKind.LOCAL + or common.is_local_dimension(dim) ): # translate sparse dimensions to tuple dtype # NOTE: the tag is the offset-provider key, and its mangled form names the diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index d95968c870..f195670437 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -577,7 +577,7 @@ def make_field( data_node: dace_nodes.AccessNode, data_type: ts.FieldType, ) -> gtir_to_sdfg_types.FieldopData: - local_dims = [dim for dim in data_type.dims if dim.kind == gtx_common.DimensionKind.LOCAL] + local_dims = [dim for dim in data_type.dims if gtx_common.is_local_dimension(dim)] if len(local_dims) == 0: # do nothing: the field domain consists of all global dimensions field_type = data_type @@ -847,7 +847,7 @@ def _make_array_shape_and_strides( neighbor_table_types = gtx_dace_args.filter_connectivity_types(self.offset_provider_type) shape = [] for dim in dims: - if dim.kind == gtx_common.DimensionKind.LOCAL: + if gtx_common.is_local_dimension(dim): # for local dimension, the size is taken from the associated connectivity type shape.append(gtx_dace_args.local_dimension_size(name, dim, neighbor_table_types)) elif gtx_dace_args.is_connectivity_identifier(name, self.offset_provider_type): @@ -935,7 +935,7 @@ def _add_storage( all_dims = gt_type.dims else: # for 'ts.ListType' use 'offset_type' as local dimension assert gt_type.dtype.offset_type is not None - assert gt_type.dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(gt_type.dtype.offset_type) assert isinstance(gt_type.dtype.element_type, ts.ScalarType) dc_dtype = gtx_dace_args.as_dace_type(gt_type.dtype.element_type) all_dims = gtx_common.order_dimensions([*gt_type.dims, gt_type.dtype.offset_type]) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 7cbe8ae87d..76cbfeb05b 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -80,7 +80,7 @@ class ValueExpr: def __post_init__(self) -> None: if isinstance(self.gt_dtype, ts.ListType): assert self.gt_dtype.offset_type is not None - assert self.gt_dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(self.gt_dtype.offset_type) @dataclasses.dataclass(frozen=True) @@ -106,7 +106,7 @@ def gt_dtype(self) -> ts.ScalarType | ts.ListType: def __post_init__(self) -> None: if isinstance(self.gt_dtype, ts.ListType): assert self.gt_dtype.offset_type is not None - assert self.gt_dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(self.gt_dtype.offset_type) @dataclasses.dataclass(frozen=True) @@ -146,7 +146,7 @@ def __post_init__(self) -> None: gtx_common.check_dims([dim for dim, _ in self.field_domain]) if isinstance(self.gt_dtype, ts.ListType): assert self.gt_dtype.offset_type is not None - assert self.gt_dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(self.gt_dtype.offset_type) assert all(dim != self.gt_dtype.offset_type for dim, _ in self.field_domain) def get_field_type(self) -> ts.FieldType: @@ -769,7 +769,7 @@ def _visit_if_branch_arg( ) # find position of the local dimension in the field layout assert isinstance(arg_desc, dace.data.Array) - assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in field_dims) + assert all(not gtx_common.is_local_dimension(dim) for dim in field_dims) extended_dims = gtx_common.order_dimensions([*field_dims, local_dim]) local_dim_pos = extended_dims.index(local_dim) inner_desc = dace.data.Array( @@ -1130,7 +1130,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: ) # The layout of connectivity tables is known. assert len(conn_type.domain) == 2 - assert conn_type.domain[1].kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(conn_type.domain[1]) conn_slice = self._construct_local_view( MemletExpr( dc_node=self.state.add_access(conn_data), @@ -1383,7 +1383,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: # The layout of connectivity tables is known. assert len(conn_type.domain) == 2 - assert conn_type.domain[1].kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(conn_type.domain[1]) conn_slice = self._construct_local_view( MemletExpr( dc_node=self.state.add_access(conn_data), diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index 0079e47638..d779e71fef 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -134,7 +134,7 @@ def _create_field_operator_impl( assert isinstance(dataflow_output_desc, dace.data.Array) assert len(dataflow_output_desc.shape) == 1 # extend the array with the local dimensions added by the field operator (e.g. `neighbors`) - assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in field_dims) + assert all(not gtx_common.is_local_dimension(dim) for dim in field_dims) assert output_edge.result.gt_dtype.offset_type is not None local_dim = output_edge.result.gt_dtype.offset_type # construct the full subset according to the canonical field domain diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py index 47baf902ec..037c2989bc 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py @@ -73,7 +73,7 @@ def get_local_view( # The invariant below is ensured by calling `make_field()` to construct `FieldopData`. # The `make_field` constructor converts any local dimension, if present, to `ListType` # element type, while leaving the field domain with all global dimensions. - assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in self.gt_type.dims) + assert all(not gtx_common.is_local_dimension(dim) for dim in self.gt_type.dims) domain_dims = [domain_range.dim for domain_range in domain] domain_indices = gtir_domain.get_element_subset( domain_dims, origin=None diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py index a2117a87a2..06c0ec272a 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py @@ -50,8 +50,9 @@ def get_map_variable(dim: gtx_common.Dimension) -> str: # fusion and map splitting rely on the names of the map variables to match the field # dimensions and decide whether two maps have the same iteration space. dim = gtx_common.as_non_staggered(dim) - suffix = "dim" if dim.kind == gtx_common.DimensionKind.LOCAL else "" - return f"i_{gtx_common.codegen_name(dim.tag)}_gtx_{dim.kind}{suffix}" + # NOTE: a local dimension has no `kind`; it keeps the name it had when `LOCAL` was a kind. + kind = "localdim" if gtx_common.is_local_dimension(dim) else str(dim.kind) + return f"i_{gtx_common.codegen_name(dim.tag)}_gtx_{kind}" def make_tasklet_connector_for(name: str) -> str: diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 6db624755c..056473db91 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -409,7 +409,7 @@ def is_local_field(type_: ts.FieldType) -> bool: Examples: >>> class V(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> is_local_field( ... ts.FieldType(dims=[V, V2E], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64)) ... ) @@ -417,7 +417,7 @@ def is_local_field(type_: ts.FieldType) -> bool: >>> is_local_field(ts.FieldType(dims=[V], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64))) False """ - return any(dim.kind == common.DimensionKind.LOCAL for dim in type_.dims) + return any(common.is_local_dimension(dim) for dim in type_.dims) def contains_local_field(type_: ts.TypeSpec) -> bool: @@ -587,7 +587,7 @@ def promote( >>> promoted.dims == [I, J, K] and promoted.dtype == dtype True - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> list_dtype = ts.ListType(element_type=dtype, offset_type=V2E) >>> promote( ... ts.FieldType(dims=[I], dtype=list_dtype), diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index efe47928ae..9fc0f7c1ae 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -184,7 +184,7 @@ class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) -class C2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2VDim(gtx.LocalDimensionIndex): ... V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py index e298515208..a510f383ce 100644 --- a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py +++ b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py @@ -42,7 +42,7 @@ class Cell(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(gtx.LocalDimensionIndex): ... C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py index bd6d613efc..e3e9a9286d 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py @@ -59,7 +59,7 @@ class Node(gtx.DimensionIndex): ... -class NeighDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class NeighDim(gtx.LocalDimensionIndex): ... def array_maker(*lists): diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py index d5406e4e08..afb2925f75 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py @@ -18,7 +18,7 @@ from gt4py.next.iterator.embedded import StridedConnectivityField -class Dummy(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class Dummy(gtx.LocalDimensionIndex): ... class LocA(gtx.DimensionIndex): ... diff --git a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py index 279a5ef023..b7c81d55fd 100644 --- a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py +++ b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py @@ -43,14 +43,14 @@ class E(gtx.DimensionIndex): ... #: N1 == N3 == N4, but N2 differs: the tag is `TaggedOffDim.tag`, the variable is `off_a`. -class TaggedOffDim(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class TaggedOffDim(gtx.LocalDimensionIndex): ... off_a = gtx.FieldOffset(TaggedOffDim.tag, source=E, target=(V, TaggedOffDim)) #: N1 == N2 == N4, but N3 differs: the local dimension is `Neigh`, the tag is `OffB`. -class Neigh(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Neigh(gtx.LocalDimensionIndex): ... OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) diff --git a/tests/next_tests/toy_connectivity.py b/tests/next_tests/toy_connectivity.py index c368b04245..531ee0867f 100644 --- a/tests/next_tests/toy_connectivity.py +++ b/tests/next_tests/toy_connectivity.py @@ -21,16 +21,16 @@ class Edge(gtx.DimensionIndex): ... class Cell(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2VDim(gtx.LocalDimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(gtx.LocalDimensionIndex): ... -class V2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2VDim(gtx.LocalDimensionIndex): ... V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) diff --git a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py index 6a47073c8a..1cbb1818c2 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py @@ -20,6 +20,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, Domain, Field, @@ -50,10 +51,10 @@ class V(DimensionIndex): ... class E(DimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class V2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2V(LocalDimensionIndex): ... class C(DimensionIndex): ... @@ -62,7 +63,7 @@ class C(DimensionIndex): ... class K(CartesianAxisIndex): ... -class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E2CO(LocalDimensionIndex): ... class A(DimensionIndex): ... @@ -77,7 +78,7 @@ class X(CartesianAxisIndex): ... class Y(CartesianAxisIndex): ... -class L(DimensionIndex, kind=DimensionKind.LOCAL): ... +class L(LocalDimensionIndex): ... class S(DimensionIndex): ... @@ -98,7 +99,7 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... class C2V(DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py index 2817dd37dc..7c8c026f58 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py @@ -21,7 +21,7 @@ class VDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... class Dim(gtx.DimensionIndex): ... -class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class LocalDim(gtx.LocalDimensionIndex): ... CartesianOffset = gtx.FieldOffset("CartesianOffset", source=Dim, target=(Dim,)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index 2b8a8a4c4f..37c650d998 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py @@ -49,7 +49,7 @@ class Edge(gtx.DimensionIndex): ... class Vertex(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) @@ -63,7 +63,7 @@ class TDim(gtx.CartesianAxisIndex): ... #: An offset whose tag differs from the name of the Python variable it is bound to, and #: from the name of its local dimension. Lowering must emit the *tag*. -class RenamedV2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class RenamedV2EDim(gtx.LocalDimensionIndex): ... renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py index 290d2914fc..2856c31149 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py @@ -19,7 +19,7 @@ """ import gt4py.next as gtx -from gt4py.next import Dims, Dimension, DimensionIndex, float64, neighbor_sum +from gt4py.next import Dims, Dimension, DimensionIndex, LocalDimensionIndex, float64, neighbor_sum from gt4py.next.ffront import source_utils from gt4py.next.ffront.source_utils import get_closure_vars_from_function @@ -30,7 +30,7 @@ class Cell(DimensionIndex): ... class Edge(DimensionIndex): ... -class C2EDim(DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py index ca5cfd4c79..afa3bba76e 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py @@ -18,6 +18,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, Field, FieldOffset, @@ -51,7 +52,7 @@ class X(CartesianAxisIndex): ... class Y(CartesianAxisIndex): ... -class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class Y2XDim(LocalDimensionIndex): ... class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... @@ -72,7 +73,7 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... class IDim(CartesianAxisIndex): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py index 1531a2b09c..aca610d25d 100644 --- a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py +++ b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py @@ -33,13 +33,13 @@ class Vertex(common.DimensionIndex): ... class Edge(common.DimensionIndex): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... -class V2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2VDim(common.LocalDimensionIndex): ... a_range = domain_utils.SymbolicRange(0, 10) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py index 869372ab78..33f6f71f43 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py @@ -30,7 +30,7 @@ class E(gtx.DimensionIndex): ... class V(gtx.DimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2VDim(gtx.LocalDimensionIndex): ... E2V = gtx.FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py index 7119c5fcb7..4be03cafb0 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py @@ -267,7 +267,7 @@ class IDim(gtx.CartesianAxisIndex): ... class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... -class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class LocalDim(gtx.LocalDimensionIndex): ... @pytest.mark.parametrize("dim, suffix", [(IDim, "ₕ"), (KDim, "ᵥ"), (LocalDim, "ₗ")]) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py index ff6890f906..abe6d9d28c 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py @@ -45,7 +45,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... float_i_field = ts.FieldType(dims=[IDim], dtype=float_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py index 250d3eabc3..5577e016be 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py @@ -17,7 +17,7 @@ from gt4py.next.type_system import type_specifications as ts -class Neighbor(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Neighbor(common.LocalDimensionIndex): ... class IDim(common.CartesianAxisIndex): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py index e4d1746ccc..ef95683b21 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py @@ -25,7 +25,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... class K(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py index 733b4f1434..ae5c1186fe 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py @@ -27,10 +27,10 @@ class dummy_neighbor(common.DimensionIndex): ... #: The local dimensions of the neighbor lists under test. Each one's `tag` is also its IR offset #: string and its offset-provider key: `UnrollReduce` looks a connectivity up by the local #: dimension of the list it reduces, so those three names must be a single string (ADR 0029). -class Dim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim(common.LocalDimensionIndex): ... -class Dim2(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim2(common.LocalDimensionIndex): ... def dummy_connectivity_type(max_neighbors: int, has_skip_values: bool): diff --git a/tests/next_tests/unit_tests/otf_tests/test_runners.py b/tests/next_tests/unit_tests/otf_tests/test_runners.py index 0a0941d792..8974d7bad1 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_runners.py +++ b/tests/next_tests/unit_tests/otf_tests/test_runners.py @@ -31,7 +31,7 @@ class Vertex(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... @pytest.fixture diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index c5daa24c9b..f307f9d5b3 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -22,6 +22,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, Domain, Infinity, @@ -58,19 +59,19 @@ class I(common.CartesianAxisIndex): ... class I_half(common.CartesianAxisIndex): ... -class C2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E(LocalDimensionIndex): ... -class V2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(LocalDimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C(LocalDimensionIndex): ... -class E2C2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2V(LocalDimensionIndex): ... class ECDim(DimensionIndex): ... @@ -634,7 +635,7 @@ def dimension_promotion_cases() -> list[ ( [[JDim, V2E], [IDim, E2C2V, KDim]], None, - "There are more than one dimension with DimensionKind 'LOCAL'.", + "There is more than one local dimension.", ), ([[JDim, V2E], [IDim, KDim]], [IDim, JDim, V2E, KDim], None), # a dimension and its staggered counterpart must not be promoted into the same field diff --git a/tests/next_tests/unit_tests/test_custom_layout_allocators.py b/tests/next_tests/unit_tests/test_custom_layout_allocators.py index 47c657c3d3..36ff69a3fe 100644 --- a/tests/next_tests/unit_tests/test_custom_layout_allocators.py +++ b/tests/next_tests/unit_tests/test_custom_layout_allocators.py @@ -29,13 +29,13 @@ class D2(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... class D0_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... -class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D1_local(common.LocalDimensionIndex): ... class D2_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... -class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D2_local(common.LocalDimensionIndex): ... class D1_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 7f975c7d2f..09291fd511 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -16,6 +16,7 @@ from gt4py._core import definitions as core_defs from gt4py.next import common from gt4py.next.common import ( + CartesianAxisIndex, DimensionIndex, DimensionKind, LocalDimensionIndex, @@ -32,7 +33,7 @@ class Vertex(DimensionIndex): ... class Edge(DimensionIndex): ... -class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... +class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... class V2E(NeighborConnectivity[Vertex, Edge], max_neighbors=4, min_neighbors=3): @@ -75,7 +76,8 @@ def test_owner_and_dimensions(self): assert V2E.Local.owner is V2E assert V2E.domain is Vertex assert V2E.codomain is Edge - assert V2E.Local.kind is DimensionKind.LOCAL + assert common.is_local_dimension(V2E.Local) + assert V2E.Local.kind is None assert issubclass(V2E.Local, DimensionIndex) def test_counts(self): @@ -85,7 +87,8 @@ def test_counts(self): def test_ownerless_local(self): assert LsqCoeff.owner is None assert (LsqCoeff.max_neighbors, LsqCoeff.min_neighbors) == (3, 3) - assert LsqCoeff.kind is DimensionKind.LOCAL + assert common.is_local_dimension(LsqCoeff) + assert LsqCoeff.kind is None def test_counts_from_local_size(self): ns = _declare( @@ -137,7 +140,7 @@ class C(NeighborConnectivity[Vertex, Edge]): ... ( """ class C(NeighborConnectivity[Vertex, Edge]): - class Local(DimensionIndex, kind=DimensionKind.LOCAL): ... + class Local(DimensionIndex): ... """, "must declare its local dimension", ), @@ -201,7 +204,7 @@ class Local(LocalDimensionIndex, size=3): ... """ class L(LocalDimensionIndex, kind=DimensionKind.HORIZONTAL): ... """, - "cannot have kind", + "is a local dimension and has no kind", ), ( """ @@ -631,3 +634,36 @@ class C(NeighborConnectivity[Vertex, Edge]): Local: typing.TypeAlias = ConstList """ ) + + +class TestLocality: + """`DimensionKind.LOCAL` is gone: localness is the class, and a local dimension has no kind.""" + + def test_no_local_kind(self): + assert set(DimensionKind.__members__) == {"HORIZONTAL", "VERTICAL"} + + @pytest.mark.parametrize( + "dim, expected", + [ + (V2E.Local, True), + (LsqCoeff, True), + (common.ConstList, True), + (Vertex, False), + (KDim, False), + ], + ) + def test_is_local_dimension(self, dim, expected): + assert common.is_local_dimension(dim) is expected + + @pytest.mark.parametrize("value", [None, "V2E", V2E, LocalDimensionIndex(0)]) + def test_is_local_dimension_of_a_non_dimension(self, value): + assert common.is_local_dimension(value) is False + + def test_display(self): + assert str(LsqCoeff) == "LsqCoeff[local]" + assert repr(LsqCoeff) == f"{LsqCoeff.tag}[local]" + + def test_layout_order_is_horizontal_local_vertical(self): + assert common.order_dimensions([KDim, LsqCoeff, Vertex]) == [Vertex, LsqCoeff, KDim] + with pytest.raises(ValueError, match="more than one local dimension"): + common.order_dimensions([Vertex, LsqCoeff, V2E.Local]) diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py index 2984713b76..f006282e53 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py @@ -14,6 +14,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, ) from gt4py.next.type_system import type_info, type_specifications as ts @@ -30,10 +31,10 @@ class JDim(CartesianAxisIndex): ... class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... class TDim(CartesianAxisIndex): ... diff --git a/typing_tests/pyright_probes.py b/typing_tests/pyright_probes.py index fbc47762a0..b89ffe6705 100644 --- a/typing_tests/pyright_probes.py +++ b/typing_tests/pyright_probes.py @@ -34,7 +34,7 @@ class Cell(gtx.DimensionIndex): ... class CellEdge(gtx.DimensionIndex): ... -class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... +class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): @@ -109,3 +109,36 @@ def shift_by_a_dimension( a: gtx.Field[gtx.Dims[KDim], gtx.float64], ) -> gtx.Field[gtx.Dims[KDim], gtx.float64]: return a(KDim + 1) + + +# -- Cartesian axis levels (ADR 0029). Each rejection carries a targeted ignore, and +# `reportUnnecessaryTypeIgnoreComment` turns any rejection that stops firing into an error. + + +def any_dimension(dim: gtx.Dimension) -> None: ... + + +def any_axis(dim: type[gtx.AnyCartesianAxisIndex]) -> None: ... + + +def declared_axis(dim: type[gtx.CartesianAxisIndex]) -> None: ... + + +def staggered_field(a: gtx.Field[gtx.Dims[Cell, gtx.Staggered[KDim]], gtx.float64]) -> None: ... + + +any_dimension(gtx.Staggered[KDim]) # a staggered dimension is still a dimension +any_dimension(V2E.Local) # local dimensions keep their place below the root +any_axis(gtx.Staggered[KDim]) +declared_axis(KDim) +_shift_staggered = gtx.Staggered[KDim] + 1 +_shift_half = KDim + 0.5 + +any_axis(Cell) # pyright: ignore[reportArgumentType] +any_axis(V2E.Local) # pyright: ignore[reportArgumentType] +declared_axis(gtx.Staggered[KDim]) # pyright: ignore[reportArgumentType] +_doubly: typing.TypeAlias = gtx.Staggered[gtx.Staggered[KDim]] # pyright: ignore[reportInvalidTypeArguments] +_location: typing.TypeAlias = gtx.Staggered[Cell] # pyright: ignore[reportInvalidTypeArguments] +_local: typing.TypeAlias = gtx.Staggered[V2E.Local] # pyright: ignore[reportInvalidTypeArguments] +_shift_location = Cell + 1 # pyright: ignore[reportOperatorIssue] +_shift_local = V2E.Local - 1 # pyright: ignore[reportOperatorIssue] diff --git a/typing_tests/pyrightconfig.json b/typing_tests/pyrightconfig.json index 4f44d9cfb8..7dc8caf8d9 100644 --- a/typing_tests/pyrightconfig.json +++ b/typing_tests/pyrightconfig.json @@ -1,5 +1,6 @@ { "typeCheckingMode": "standard", "reportMissingImports": "error", - "reportMissingTypeStubs": "none" + "reportMissingTypeStubs": "none", + "reportUnnecessaryTypeIgnoreComment": "error" } From e565d9c398c1754c690e12f12be5b3af28dc4c4e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 00:42:34 +0200 Subject: [PATCH 13/13] test[next]: a local dimension cannot be staggered --- tests/next_tests/unit_tests/test_neighbor_connectivity.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 09291fd511..8e47de451f 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -659,6 +659,12 @@ def test_is_local_dimension(self, dim, expected): def test_is_local_dimension_of_a_non_dimension(self, value): assert common.is_local_dimension(value) is False + @pytest.mark.parametrize("local", [V2E.Local, LsqCoeff, LocalDimensionIndex]) + def test_a_local_dimension_cannot_be_staggered(self, local): + # what keeps `is_local_dimension` total: there is no staggered local dimension + with pytest.raises(TypeError, match="not a declared Cartesian axis"): + common.Staggered[local] + def test_display(self): assert str(LsqCoeff) == "LsqCoeff[local]" assert repr(LsqCoeff) == f"{LsqCoeff.tag}[local]"