From c215b52201fefe58a7fe2ec87c1bf1e794c4dce5 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/17] 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 | 1 + 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, 936 insertions(+), 23 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 9c93fd9aa0..e663e462f3 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -23,6 +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) ### Frontend and Parsing #frontend diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index e665024d7d..8bf304946f 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -31,6 +31,8 @@ Domain, Field, GridType, + LocalDimensionIndex, + NeighborConnectivity, Staggered, UnitRange, as_non_staggered, @@ -120,6 +122,8 @@ "Dimension", "DimensionIndex", "DimensionKind", + "LocalDimensionIndex", + "NeighborConnectivity", "Staggered", "resolve", "Dims", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 89014c25c0..ed4047d84e 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 @@ -1027,7 +1028,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: ... @@ -1035,8 +1038,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 @@ -1557,8 +1560,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() @@ -1731,6 +1734,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." @@ -1878,3 +1883,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 0e3aaeab8b..f8c422f134 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -239,7 +239,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). @@ -314,7 +316,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() @@ -366,8 +371,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 3d09d7ecb4..21ab0cfdcd 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 1b6cefc6b2..131647baf4 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 4f3ddd732c..7b9e025f83 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -154,7 +154,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] @@ -171,8 +176,10 @@ def as_scalar(self) -> xtyping.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() @@ -1152,8 +1159,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() @@ -1293,8 +1302,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 20ed020333..67f3b7e9db 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -275,3 +275,72 @@ main: | import xarray a: xarray.NamedArray + + - 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 ba01559731e232b79a15da8b5975e6fbf8807102 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/17] 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 ed4047d84e..ad6fc82e1e 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1931,13 +1931,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) @@ -1973,6 +1973,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) @@ -1995,6 +2000,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) @@ -2011,7 +2021,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): ... @@ -2108,6 +2119,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). @@ -2132,7 +2146,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( @@ -2141,6 +2155,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 21ab0cfdcd..85eaafe610 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 b762da1ec9..998ba43238 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 889f44d37a208c43ff57a985d5cd576fe0d0c596 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/17] 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 ad6fc82e1e..f8bb5d4233 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1958,6 +1958,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;" @@ -1993,21 +2015,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 @@ -2077,14 +2091,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 207dae3f88482a2de25759d2ab8552e936aa51a0 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/17] 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 64f52fdb8b..89a743ce18 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -64,6 +64,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 8bf304946f..94690b7d0e 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -39,6 +39,7 @@ domain, flip_staggered, is_staggered, + local_dimension_of, resolve, unit_range, ) @@ -136,6 +137,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 f8bb5d4233..b1b53e28db 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1911,6 +1911,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, @@ -1928,7 +1930,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]: @@ -1949,7 +1952,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 @@ -1973,12 +1979,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( @@ -2049,9 +2055,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] @@ -2088,12 +2093,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 ( @@ -2110,14 +2120,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 @@ -2133,6 +2148,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, @@ -2152,7 +2181,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 998ba43238..0ba94993a1 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 67f3b7e9db..45ddb1037f 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -335,7 +335,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 @@ -343,4 +345,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 d024752ae3..d3c356efda 100644 --- a/uv.lock +++ b/uv.lock @@ -1441,6 +1441,7 @@ typing = [ ] typing-exports = [ { name = "mypy", extra = ["faster-cache"] }, + { name = "pyright" }, { name = "pytest-mypy-plugins" }, { name = "types-decorator" }, { name = "types-docutils" }, @@ -1596,6 +1597,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" }, @@ -3124,6 +3126,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 bf5c95047f47e72a71ac8aaac870f0874c18805a 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/17] 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 b1b53e28db..4e9ab67196 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1222,7 +1222,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 809137d4e98133e37cb2bbbf288e8a59e4eaca21 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 00:22:58 +0200 Subject: [PATCH 06/17] feat[next]: declare connectivities as classes; FieldOffset derived from them --- docs/user/next/QuickstartGuide.md | 44 +++++++------- .../exercises/2_divergence_exercise.ipynb | 2 +- .../2_divergence_exercise_solution.ipynb | 2 +- .../exercises/3_gradient_exercise.ipynb | 2 +- .../3_gradient_exercise_solution.ipynb | 2 +- .../workshop/exercises/4_curl_exercise.ipynb | 2 +- .../exercises/4_curl_exercise_solution.ipynb | 2 +- .../exercises/5_vector_laplace_exercise.ipynb | 8 +-- .../5_vector_laplace_exercise_solution.ipynb | 8 +-- .../8_diffusion_exercise_solution.ipynb | 2 +- docs/user/next/workshop/exercises/helpers.py | 33 +++++++---- docs/user/next/workshop/slides/slides_2.ipynb | 5 +- src/gt4py/next/common.py | 58 ++++++++++++------- src/gt4py/next/constructors.py | 2 +- src/gt4py/next/ffront/fbuiltins.py | 17 +++++- src/gt4py/next/iterator/embedded.py | 21 +++++-- src/gt4py/next/iterator/tracing.py | 2 + .../iterator/type_system/type_synthesizer.py | 2 +- src/gt4py/next/type_system/type_info.py | 4 +- .../integration_tests/cases_utils.py | 36 +++++++----- ..._write_back_buffer_elimination_lowering.py | 5 +- .../ffront_tests/test_compiled_program.py | 4 +- .../iterator_tests/test_builtins.py | 2 +- .../test_strided_offset_provider.py | 2 +- .../multi_feature_tests/fvm_nabla_setup.py | 6 +- .../test_offset_dimensions_names.py | 4 +- tests/next_tests/toy_connectivity.py | 20 ++++--- .../embedded_tests/test_nd_array_field.py | 11 ++-- .../test_decorator_domain_deduction.py | 7 ++- .../ffront_tests/test_foast_to_gtir.py | 7 ++- .../ffront_tests/test_source_utils.py | 7 ++- .../ffront_tests/test_type_deduction.py | 17 ++++-- .../ir_utils_tests/test_domain_utils.py | 6 +- .../test_embedded_field_with_list.py | 5 +- .../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 +- .../dace_tests/test_dace_translation.py | 2 +- tests/next_tests/unit_tests/test_common.py | 11 ++-- .../test_custom_layout_allocators.py | 4 +- .../unit_tests/test_neighbor_connectivity.py | 29 +++++++++- .../type_system_tests/test_type_info.py | 5 +- 44 files changed, 262 insertions(+), 158 deletions(-) diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 1e1cbfc280..08c37acb03 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -218,28 +218,28 @@ edge_values = gtx.as_field([EdgeDim], np.zeros((12,))) +++ -You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _field offset_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements. +You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _connectivity_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements. To understand this transform, you can look at the edge-to-cell connectivity table `edge_to_cell_table` listed above. This table has the same shape as the output of the transform, that is, one dimension over the edges and another _local_ dimension. The table stores indices into a field over cells, the transform essentially gives you another field where the indices have been replaced with the values in the cell field at the corresponding indices. Another way to look at it is that transform uses the edge-to-cell connectivity to look up all the cell neighbors of edges, and associates the values of those neighbor cells with each edge. -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: +You can use the connectivity `E2C` declared below to transform a field over cells to a field over edges using the edge-to-cell connectivities. It is declared as a class: for each edge (`EdgeDim`), a list of neighbor cells (`CellDim`), indexed by its nested local dimension `E2C.Local`: ```{code-cell} ipython3 -class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... -E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim)) +class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + class Local(gtx.LocalDimensionIndex): ... ``` -The field offset is named by its local dimension's `tag`, and the offset provider below is keyed by the same `tag`, so all three refer to one connectivity. Note that the field offset does not contain the actual connectivity table, that's provided through an _offset provider_: +Note that the declaration does not contain the actual connectivity table, that's provided through an _offset provider_, keyed by the local dimension's `tag`: ```{code-cell} ipython3 -E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2CDim], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) +E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2C.Local], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) ``` The field operator `nearest_cell_to_edge` below shows an example of applying this transform. There is a little twist though: the subscript in `E2C[0]` means that only the value of the first connected cell is taken, the second (if exists) is ignored. -Pay attention to the syntax where the field offset `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. +Pay attention to the syntax where the connectivity `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. ```{code-cell} ipython3 @gtx.field_operator @@ -250,7 +250,7 @@ def nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64]) -> gtx. def run_nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): nearest_cell_to_edge(cell_values, out=out) -run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider}) print("0th adjacent cell's value: {}".format(edge_values.asnumpy())) ``` @@ -265,19 +265,19 @@ Running the above snippet results in the following edge field: #### Using reductions on connected mesh elements -Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2CDim) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2CDim` dimension. +Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2C.Local) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2C.Local` dimension. ```{code-cell} ipython3 @gtx.field_operator def sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64]) -> gtx.Field[Dims[EdgeDim], float64]: - # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2CDim], float64] - return neighbor_sum(cells(E2C), axis=E2CDim) + # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2C.Local], float64] + return neighbor_sum(cells(E2C), axis=E2C.Local) @gtx.program def run_sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): sum_adjacent_cells(cells, out=out) -run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider}) print("sum of adjacent cells: {}".format(edge_values.asnumpy())) ``` @@ -376,13 +376,13 @@ print("where nested tuple return: {}".format(((result_1.asnumpy(), result_2.asnu #### Implementing the pseudo-laplacian -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: +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 declare the connectivity 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): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=EdgeDim, target=(CellDim, C2EDim)) +class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + class Local(gtx.LocalDimensionIndex): ... -C2E_offset_provider = gtx.as_connectivity([CellDim, C2EDim], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) +C2E_offset_provider = gtx.as_connectivity([CellDim, C2E.Local], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) ``` **Weights of edge differences:** @@ -410,7 +410,7 @@ edge_weights = np.array([ [0, -1, -1], # cell 5 ], dtype=np.float64) -edge_weight_field = gtx.as_field([CellDim, C2EDim], edge_weights) +edge_weight_field = gtx.as_field([CellDim, C2E.Local], edge_weights) ``` Now you have everything to implement the pseudo-laplacian. Its field operator requires the cell field and the edge weights as inputs, and outputs a cell field of the same shape as the input. @@ -422,9 +422,9 @@ The second lines first creates a temporary field using `edge_differences(C2E)`, ```{code-cell} ipython3 @gtx.field_operator def pseudo_lap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: edges = cells(E2C[0]) # type: gtx.Field[Dims[EdgeDim], float64] - return neighbor_sum(edges(C2E) * edge_weights, axis=C2EDim) + return neighbor_sum(edges(C2E) * edge_weights, axis=C2E.Local) ``` The program itself is just a shallow wrapper over the `pseudo_lap` field operator. The significant part is how offset providers for both the edge-to-cell and cell-to-edge connectivities are supplied when the program is called: @@ -432,7 +432,7 @@ The program itself is just a shallow wrapper over the `pseudo_lap` field operato ```{code-cell} ipython3 @gtx.program def run_pseudo_laplacian(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64], + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64], out : gtx.Field[Dims[CellDim], float64]): pseudo_lap(cells, edge_weights, out=out) @@ -441,7 +441,7 @@ result_pseudo_lap = gtx.as_field([CellDim], np.zeros(shape=(6,))) run_pseudo_laplacian(cell_values, edge_weight_field, result_pseudo_lap, - offset_provider={E2CDim.tag: E2C_offset_provider, C2EDim.tag: C2E_offset_provider}) + offset_provider={E2C.Local.tag: E2C_offset_provider, C2E.Local.tag: C2E_offset_provider}) print("pseudo-laplacian: {}".format(result_pseudo_lap.asnumpy())) ``` @@ -451,7 +451,7 @@ As a closure, here is an example of chaining field operators, which is very simp ```{code-cell} ipython3 @gtx.field_operator def pseudo_laplap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: return pseudo_lap(pseudo_lap(cells, edge_weights), edge_weights) ``` diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb index b0a1980d0f..21bf2d25d8 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb @@ -126,7 +126,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb index 573ee6a44e..86c8d33ac7 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb @@ -131,7 +131,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb index 2b422b1823..fb2282ab22 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb @@ -123,7 +123,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb index 85044b989f..43196507ce 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb @@ -136,7 +136,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb index dc321f1bdd..b99c6f6d3f 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb @@ -147,7 +147,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E.Local.tag: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb index 251fe8239a..de040ccb93 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb @@ -152,7 +152,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E.Local.tag: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb index 30f568de6f..174699e350 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb @@ -293,10 +293,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E.Local.tag: c2e_connectivity,\n", + " V2E.Local.tag: v2e_connectivity,\n", + " E2V.Local.tag: e2v_connectivity,\n", + " E2C.Local.tag: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb index eaeb8c7b02..81836edbd5 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb @@ -314,10 +314,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E.Local.tag: c2e_connectivity,\n", + " V2E.Local.tag: v2e_connectivity,\n", + " E2V.Local.tag: e2v_connectivity,\n", + " E2C.Local.tag: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb index b278cee26d..edd65ac9e5 100644 --- a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb @@ -169,7 +169,7 @@ " kappa,\n", " dt,\n", " out=(divergence_gt4py_1, divergence_gt4py_2),\n", - " offset_provider={E2C2V.value: e2c2v_connectivity, V2E.value: v2e_connectivity},\n", + " offset_provider={E2C2V.Local.tag: e2c2v_connectivity, V2E.Local.tag: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py_1.asnumpy(), divergence_ref_1)\n", diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index c398524538..07fb984328 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -11,7 +11,13 @@ import gt4py.next as gtx from gt4py.next.iterator.embedded import MutableLocatedField from gt4py.next import neighbor_sum, where, Dims -from gt4py.next import Dimension, DimensionIndex, DimensionKind, FieldOffset +from gt4py.next import ( + Dimension, + DimensionIndex, + DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, +) from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( run_gtfn as gtfn_cpu, @@ -389,31 +395,36 @@ class E(DimensionIndex): ... class K(DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E(NeighborConnectivity[C, E]): + class Local(LocalDimensionIndex): ... -C2E = FieldOffset(C2EDim.tag, source=E, target=(C, C2EDim)) +C2EDim = C2E.Local -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(NeighborConnectivity[V, E]): + class Local(LocalDimensionIndex): ... -V2E = FieldOffset(V2EDim.tag, source=E, target=(V, V2EDim)) +V2EDim = V2E.Local -class E2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(NeighborConnectivity[E, V]): + class Local(LocalDimensionIndex): ... -E2V = FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) +E2VDim = E2V.Local -class E2CDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C(NeighborConnectivity[E, C]): + class Local(LocalDimensionIndex): ... -E2C = FieldOffset(E2CDim.tag, source=C, target=(E, E2CDim)) +E2CDim = E2C.Local -class E2C2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2V(NeighborConnectivity[E, V]): + class Local(LocalDimensionIndex): ... -E2C2V = FieldOffset(E2C2VDim.tag, source=V, target=(E, E2C2VDim)) +E2C2VDim = E2C2V.Local diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index db8f370abc..f16f1560bc 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -273,10 +273,11 @@ "metadata": {}, "outputs": [], "source": [ - "class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...\n", + "class E2C(gtx.NeighborConnectivity[Edge, Cell]):\n", + " class Local(gtx.LocalDimensionIndex): ...\n", "\n", "\n", - "E2C = gtx.FieldOffset(E2CDim.tag, source=Cell, target=(Edge, E2CDim))" + "E2CDim = E2C.Local" ] }, { diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 4e9ab67196..00bf075417 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -284,6 +284,14 @@ def __init_subclass__(cls, /, kind: Optional[DimensionKind] = None, **kwargs: An ) if kind is not None: cls.kind = kind + if cls.kind is DimensionKind.LOCAL and not any( + "_local_dimension_root" in base.__dict__ for base in cls.__mro__ + ): + raise TypeError( + f"'{cls.__qualname__}': a local dimension is declared by subclassing" + " 'LocalDimensionIndex', or as the nested 'Local' class of a" + " 'NeighborConnectivity', not with 'kind=DimensionKind.LOCAL'." + ) def __init__(self, value: int) -> None: self.value = value @@ -1633,8 +1641,8 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: >>> class I(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... >>> class J(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... >>> class K(DimensionIndex, 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]) @@ -1799,24 +1807,6 @@ def __init_subclass__(cls, /, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) -class ConstListDim(DimensionIndex, kind=DimensionKind.LOCAL): - """ - The local dimension of a list whose length is known at compile time (`make_const_list`). - - Declared here, once, because it must be a *single* class. It used to be built - independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless - while dimensions compared by `(name, kind)` -- the two instances were equal. Under - nominal identity (ADR 0028) two declarations would be two different dimensions, and the - `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s - built by embedded execution. - - TODO: becomes an owner-less local dimension with an explicit size, generalising this from - length 1 to length *n*, once local dimensions know their connectivity. - """ - - __slots__ = () - - def _reduce_staggered(cls: StaggeredMeta) -> Any: """ Pickle a staggered dimension through its base, falling back to by-reference. @@ -1888,7 +1878,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. @@ -1906,6 +1896,9 @@ class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): __slots__ = () + kind: ClassVar[DimensionKind] = DimensionKind.LOCAL + _local_dimension_root: ClassVar[bool] = True + #: 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 @@ -1947,6 +1940,24 @@ def _check_neighbor_count(cls: type, name: str, count: Optional[int]) -> Optiona return int(count) +class ConstListDim(LocalDimensionIndex): + """ + The local dimension of a list whose length is known at compile time (`make_const_list`). + + Declared here, once, because it must be a *single* class. It used to be built + independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless + while dimensions compared by `(name, kind)` -- the two instances were equal. Under + nominal identity (ADR 0028) two declarations would be two different dimensions, and the + `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s + built by embedded execution. + + TODO: becomes an owner-less local dimension with an explicit size, generalising this from + length 1 to length *n*, once local dimensions know their connectivity. + """ + + __slots__ = () + + class ConnectivityMeta(type): """ Metaclass of `NeighborConnectivity` declarations. @@ -2030,7 +2041,10 @@ 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.origin, cls._local()), + _derived=True, ) type.__setattr__(cls, "_field_offset", field_offset) return field_offset diff --git a/src/gt4py/next/constructors.py b/src/gt4py/next/constructors.py index f1434bc954..ba6c1e5cab 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/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 942a5dcc10..0916417a5f 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -11,6 +11,7 @@ import inspect import math import operator +import warnings from builtins import bool, float, int, tuple # noqa: A004 shadowing a Python built-in from types import UnionType from typing import ( @@ -484,14 +485,26 @@ class FieldOffset(runtime.Offset): value: str source: common.Dimension target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] + #: Set when derived from a `NeighborConnectivity` declaration, which is not deprecated. + _derived: bool = dataclasses.field(default=False, repr=False, compare=False, kw_only=True) @functools.cached_property def _cache(self) -> dict: return {} 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.") + if len(self.target) == 2: + if self.target[1].kind != common.DimensionKind.LOCAL: + raise ValueError("Second dimension in offset must be a local dimension.") + if not self._derived: + warnings.warn( + "Declaring an unstructured connectivity with 'FieldOffset' is deprecated;" + " declare a 'NeighborConnectivity' class instead (see ADR 0029):\n" + " class V2E(gtx.NeighborConnectivity[Vertex, Edge]):\n" + " class Local(gtx.LocalDimensionIndex): ...", + DeprecationWarning, + stacklevel=3, + ) def __gt_type__(self) -> ts.OffsetType: return ts.OffsetType(source=self.source, target=self.target, tag=self.value) diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 7b9e025f83..3401252252 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -1400,10 +1400,10 @@ def constant_field(value: Any, dtype_like: Optional[core_defs.DTypeLike] = None) @builtins.shift.register(EMBEDDED) def shift( - *offsets: Union[runtime.Offset, OffsetPart], + *offsets: Union[runtime.Offset, type[common.NeighborConnectivity], OffsetPart], ) -> Callable[[ItIterator], ItIterator]: def impl(it: ItIterator) -> ItIterator: - return it.shift(*list(o.value if isinstance(o, runtime.Offset) else o for o in offsets)) + return it.shift(*list(_as_offset_tag(o) for o in offsets)) return impl @@ -1448,9 +1448,20 @@ def __gt_type__(self) -> ts.ListType: ) +def _as_offset_tag( + offset: runtime.Offset | type[common.NeighborConnectivity] | OffsetPart, +) -> OffsetPart: + if isinstance(offset, common.ConnectivityMeta): + return offset.Local.tag + return offset.value if isinstance(offset, runtime.Offset) else offset + + @builtins.neighbors.register(EMBEDDED) -def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: - offset_str = offset.value if isinstance(offset, runtime.Offset) else offset +def neighbors(offset: runtime.Offset | type[common.NeighborConnectivity], it: ItIterator) -> _List: + field_offset: runtime.Offset = ( + offset.__gt_field_offset__() if isinstance(offset, common.ConnectivityMeta) else offset + ) + offset_str = _as_offset_tag(field_offset) assert isinstance(offset_str, str) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None @@ -1462,7 +1473,7 @@ def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: for i in range(connectivity.__gt_type__().max_neighbors) if (shifted := it.shift(offset_str, i)).can_deref() ), - offset=offset, + offset=field_offset, ) diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 4e5b8d4b5a..69ab9676fd 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -153,6 +153,8 @@ def make_node(o): # it, see `execute_shift`); decide whether to fold it into the shift value or forbid it. assert o.offset == 0 return im.cartesian_offset(o.domain_dim, o.codomain) + if isinstance(o, common.ConnectivityMeta): + return OffsetLiteral(value=o.Local.tag) if callable(o): if o.__name__ == "": return lambdadef(o) diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 7415bcb7a1..17b0949c6a 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, diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 2ce4672a77..da0cf84b1a 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -408,7 +408,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)) ... ) @@ -586,7 +586,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 464bf7fa1f..2efc73b47c 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -35,7 +35,17 @@ # Both modules used to declare their own `Dimension("Vertex")` etc., which compared equal; under # nominal identity (ADR 0028) that would be two different dimensions, and tests that mix a # `toy_connectivity` connectivity with a `cases_utils` mesh would silently stop matching. -from next_tests.toy_connectivity import C2EDim, Cell, E2VDim, Edge, V2EDim, Vertex +from next_tests.toy_connectivity import ( + C2E, + C2EDim, + Cell, + E2V, + E2VDim, + Edge, + V2E, + V2EDim, + Vertex, +) __all__ = [ @@ -184,13 +194,11 @@ class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) -class C2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2V(gtx.NeighborConnectivity[Cell, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) -C2V = gtx.FieldOffset(C2VDim.tag, source=Vertex, target=(Cell, C2VDim)) +C2VDim = C2V.Local size = 10 @@ -308,28 +316,28 @@ def simple_mesh(allocator) -> MeshDescriptor: e2v_arr = np.asarray(e2v_arr, dtype=gtx.IndexType) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E.Local.tag: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 4}, codomain=Edge, data=v2e_arr, skip_value=None, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V.Local.tag: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V.Local.tag: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 4}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E.Local.tag: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 4}, codomain=Edge, data=c2e_arr, @@ -403,28 +411,28 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E.Local.tag: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 5}, codomain=Edge, data=v2e_arr, skip_value=common._DEFAULT_SKIP_VALUE, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V.Local.tag: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V.Local.tag: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 3}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E.Local.tag: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 3}, codomain=Edge, data=c2e_arr, 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 fa65a38b40..f03f970507 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,10 +42,11 @@ class Cell(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2EDim = C2E.Local C2E_TABLE = np.array( [ diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py index 4a642f5cc4..57a2af8dfc 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py @@ -252,7 +252,7 @@ def test_compile_unstructured(unstructured_case, compile_testee_unstructured): compile_testee_unstructured(*args, offset_provider=unstructured_case.offset_provider, **kwargs) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E.Local.tag].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), @@ -317,7 +317,7 @@ def test_compile_unstructured_for_two_offset_providers( *args, offset_provider=unstructured_case.offset_provider, **kwargs ) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E.Local.tag].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), 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 b16461447a..094686be80 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 bd675b5f51..626e83d075 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 @@ -17,7 +17,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/integration_tests/multi_feature_tests/fvm_nabla_setup.py b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py index 4fbb01c72d..7051170dd7 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py @@ -39,11 +39,7 @@ # NOTE: imported, not redeclared. Under nominal identity (ADR 0028) a same-named declaration # here would be a different dimension from the one `toy_connectivity` declares, where the old # `Dimension("...")` values compared equal -- and tests mix objects from both modules. -from next_tests.toy_connectivity import E2VDim, Edge, V2EDim, Vertex - - -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) +from next_tests.toy_connectivity import E2V, E2VDim, Edge, V2E, V2EDim, Vertex def assert_close(expected, actual): 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..3cf6519d27 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..28fffdc732 100644 --- a/tests/next_tests/toy_connectivity.py +++ b/tests/next_tests/toy_connectivity.py @@ -21,22 +21,26 @@ class Edge(gtx.DimensionIndex): ... class Cell(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2V(gtx.NeighborConnectivity[Edge, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class V2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2V(gtx.NeighborConnectivity[Vertex, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) -V2V = gtx.FieldOffset(V2VDim.tag, source=Vertex, target=(Vertex, V2VDim)) +V2EDim = V2E.Local +E2VDim = E2V.Local +C2EDim = C2E.Local +V2VDim = V2V.Local # 3x3 periodic edges cells # 0 - 1 - 2 - 0 1 2 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 bef6a29c8a..3acc165129 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, DimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Field, DimensionIndex, @@ -49,10 +50,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): ... @@ -61,7 +62,7 @@ class C(DimensionIndex): ... class K(DimensionIndex): ... -class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E2CO(LocalDimensionIndex): ... class A(DimensionIndex): ... @@ -76,7 +77,7 @@ class X(DimensionIndex): ... class Y(DimensionIndex): ... -class L(DimensionIndex, kind=DimensionKind.LOCAL): ... +class L(LocalDimensionIndex): ... class S(DimensionIndex): ... @@ -97,7 +98,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 25284281ef..17a19ff744 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,11 +21,14 @@ class VDim(gtx.DimensionIndex, 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,)) -UnstructuredOffset = gtx.FieldOffset(LocalDim.tag, source=Dim, target=(Dim, LocalDim)) + + +class UnstructuredOffset(gtx.NeighborConnectivity[Dim, Dim]): + Local = LocalDim def test_domain_deduction_cartesian(): 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 f696daf1b4..65a421c171 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,10 +49,11 @@ class Edge(gtx.DimensionIndex): ... class Vertex(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) +V2EDim = V2E.Local class TDim(gtx.DimensionIndex): ... @@ -63,7 +64,7 @@ class TDim(gtx.DimensionIndex): ... #: 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..7051400ed8 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,10 +30,11 @@ class Cell(DimensionIndex): ... class Edge(DimensionIndex): ... -class C2EDim(DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2EDim = C2E.Local CField = gtx.Field[Dims[Cell], float64] EField = gtx.Field[Dims[Edge], float64] 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 88a8c640d8..c21b88be36 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,8 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, Field, FieldOffset, astype, @@ -50,7 +52,11 @@ class X(DimensionIndex): ... class Y(DimensionIndex): ... -class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class Y2X(NeighborConnectivity[Y, X]): + class Local(LocalDimensionIndex): ... + + +Y2XDim = Y2X.Local class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... @@ -71,7 +77,11 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + + +V2EDim = V2E.Local class IDim(DimensionIndex): ... @@ -286,7 +296,6 @@ def domain_comparison(a: Field[[TDim], float], b: Field[[TDim], float]): @pytest.fixture def premap_setup(): - Y2X = FieldOffset(Y2XDim.tag, source=X, target=(Y, Y2XDim)) return X, Y, Y2XDim, Y2X @@ -567,8 +576,6 @@ def as_offset_dtype(a: Field[[ADim, BDim], float], b: Field[[BDim], float]): def test_as_offset_non_cartesian(): - V2E = FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) - def as_offset_neighbor(a: Field[[Edge], float], b: Field[[Edge], int]): return a(as_offset(V2E, b)) 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 7120882c75..793f802c9f 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 12c0b75649..be0048c285 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 @@ -29,10 +29,11 @@ class E(gtx.DimensionIndex): ... class V(gtx.DimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2V(gtx.NeighborConnectivity[E, V]): + class Local(gtx.LocalDimensionIndex): ... -E2V = gtx.FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) +E2VDim = E2V.Local # 0 --0-- 1 --1-- 2 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 eafe45b9d3..23c4dbaf07 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 5bceba0e53..7dfbf60934 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.DimensionIndex): ... 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 3f6cd212ae..d58b728e10 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.DimensionIndex, 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 5451a2dc44..3f93017c67 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 0028). -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 2ea966ff7e..f75471fcda 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/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index 1bec200fad..873a0f5bb4 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -94,7 +94,7 @@ def test_find_constant_symbols(has_unit_stride, disable_field_origin): itir.SetAt( expr=im.as_fieldop( im.lambda_("it")(im.reduce("plus", im.literal_from_value(1.0))(im.deref("it"))) - )(im.as_fieldop_neighbors(V2E.value, "x")), + )(im.as_fieldop_neighbors(V2E.Local.tag, "x")), domain=im.get_field_domain(gtx_common.GridType.UNSTRUCTURED, "y", VFTYPE.dims), target=itir.SymRef(id="y"), ) diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index 500a355dbb..d1b9f97560 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -22,6 +22,7 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Infinity, UnitRange, @@ -57,19 +58,19 @@ class I(common.DimensionIndex): ... class I_half(common.DimensionIndex): ... -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): ... 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 f5257506fb..53312af3ab 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.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class D0_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... -class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D1_local(common.LocalDimensionIndex): ... class D2_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... -class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D2_local(common.LocalDimensionIndex): ... class D1_vertical(common.DimensionIndex, 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 5d09dd8d76..16a390e112 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -137,7 +137,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", ), @@ -202,6 +202,12 @@ class L(LocalDimensionIndex, size=1.5): ... """, "must be an integer", ), + ( + """ + class L(DimensionIndex, kind=DimensionKind.LOCAL): ... + """, + "subclassing 'LocalDimensionIndex'", + ), ], ) def test_rejected(self, source, match): @@ -473,3 +479,24 @@ 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) + + +class TestFieldOffsetDeprecation: + def test_unstructured_field_offset_warns(self): + from gt4py.next import FieldOffset + + with pytest.warns(DeprecationWarning, match="NeighborConnectivity"): + FieldOffset(V2E.Local.tag, source=Edge, target=(Vertex, V2E.Local)) + + def test_derived_and_cartesian_field_offsets_do_not_warn(self, recwarn): + from gt4py.next import FieldOffset + + class_ns = _declare( + """ + class C2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + """ + ) + class_ns["C2E"].__gt_field_offset__() + FieldOffset("Koff", source=KDim, target=(KDim,)) + assert not [w for w in recwarn if issubclass(w.category, DeprecationWarning)] 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 6ee13a358c..07f36a5737 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, DimensionIndex, DimensionKind, + LocalDimensionIndex, ) from gt4py.next.type_system import type_info, type_specifications as ts from gt4py.next.ffront import type_specifications as ts_ffront @@ -29,10 +30,10 @@ class JDim(DimensionIndex): ... class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... class TDim(DimensionIndex): ... From 7f1b5c87e7933c11436c9ed28c83efa1512c78b7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 15:45:40 +0200 Subject: [PATCH 07/17] fix[next]: review fixes for the connectivity declarations in the tree - write an adopted or shared 'Local' as 'Local: TypeAlias = ...', the spelling that keeps it a type for mypy as well as pyright - refuse to adopt the 'make_const_list' local dimension, which now is one - drop a stale commented-out string-keyed offset provider --- src/gt4py/next/common.py | 5 +++++ src/gt4py/next/iterator/embedded.py | 2 +- src/gt4py/next/iterator/tracing.py | 2 +- .../ffront_tests/test_neighbor_connectivity.py | 4 +++- .../feature_tests/ffront_tests/test_program.py | 2 +- .../ffront_tests/test_decorator_domain_deduction.py | 4 +++- .../unit_tests/test_neighbor_connectivity.py | 10 ++++++++++ 7 files changed, 24 insertions(+), 5 deletions(-) diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 00bf075417..d91bbff004 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -2114,6 +2114,11 @@ def __init_subclass__( " ('class Local(LocalDimensionIndex): ...') or by adopting one" " ('Local: TypeAlias = SomeLocalDim')." ) + if local is ConstListDim: + raise TypeError( + f"'{name}' cannot adopt '{ConstListDim.__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/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 3401252252..5a9e4011cc 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -1452,7 +1452,7 @@ def _as_offset_tag( offset: runtime.Offset | type[common.NeighborConnectivity] | OffsetPart, ) -> OffsetPart: if isinstance(offset, common.ConnectivityMeta): - return offset.Local.tag + return common.local_dimension_of(offset).tag return offset.value if isinstance(offset, runtime.Offset) else offset diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 69ab9676fd..094c85925f 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -154,7 +154,7 @@ def make_node(o): assert o.offset == 0 return im.cartesian_offset(o.domain_dim, o.codomain) if isinstance(o, common.ConnectivityMeta): - return OffsetLiteral(value=o.Local.tag) + return OffsetLiteral(value=common.local_dimension_of(o).tag) if callable(o): if o.__name__ == "": return lambdadef(o) 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..1b92138e0e 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,8 @@ """A `NeighborConnectivity` declaration used directly in DSL code, on every backend.""" +import typing + import numpy as np import pytest @@ -33,7 +35,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 diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py index d882a88ec4..abad348c92 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py @@ -60,7 +60,7 @@ def shift_by_one(in_field: cases.IFloatField) -> cases.IFloatField: # direct call to field operator # TODO(tehrengruber): slicing located fields not supported currently - # shift_by_one(in_field, out=out_field[:-1], offset_provider={"Ioff": IDim}) + # shift_by_one(in_field, out=out_field[:-1]) @gtx.program def shift_by_one_program(in_field: cases.IFloatField, out_field: cases.IFloatField): 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 17a19ff744..9cdb145fca 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 @@ -6,6 +6,8 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import typing + import pytest import gt4py.next as gtx @@ -28,7 +30,7 @@ class LocalDim(gtx.LocalDimensionIndex): ... class UnstructuredOffset(gtx.NeighborConnectivity[Dim, Dim]): - Local = LocalDim + Local: typing.TypeAlias = LocalDim def test_domain_deduction_cartesian(): diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 16a390e112..86bb41f1a3 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -500,3 +500,13 @@ class Local(LocalDimensionIndex): ... class_ns["C2E"].__gt_field_offset__() FieldOffset("Koff", source=KDim, target=(KDim,)) assert not [w for w in recwarn if issubclass(w.category, DeprecationWarning)] + + +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 = ConstListDim + """ + ) From 362ac09ce374363cd0600979ca7edb4b863a62f9 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 08/17] 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 | 38 ++++++++++++++++-- .../test_offset_dimensions_names.py | 24 ++++------- .../transforms_tests/test_unroll_reduce.py | 8 ++-- 13 files changed, 171 insertions(+), 86 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 89a743ce18..3a5ea4af72 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -263,8 +263,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 d91bbff004..083c4d4359 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1471,6 +1471,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 f8c422f134..1bbdf3ef82 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -988,8 +988,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 5a9e4011cc..73d7d00231 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -578,7 +578,12 @@ def execute_shift( if tag == _CONST_DIM.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] @@ -1009,7 +1014,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[ @@ -1420,16 +1425,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) @@ -1543,7 +1558,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( @@ -1553,7 +1573,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: @@ -1800,7 +1820,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 ca67fb18af..36a80cb97c 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -91,7 +91,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 924758c727..2ac65d0e58 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 @@ -1321,7 +1321,9 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if offset_type == _CONST_DIM: # 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 @@ -1375,7 +1377,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 @@ -1477,7 +1481,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" @@ -1489,7 +1495,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 1b92138e0e..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,7 @@ """A `NeighborConnectivity` declaration used directly in DSL code, on every backend.""" +import dataclasses import typing import numpy as np @@ -120,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]: @@ -130,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( @@ -145,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 3cf6519d27..df066e21cf 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.LocalDimensionIndex): ... 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 3f93017c67..7c81aa5514 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 22b8382c878efdd51b990ca0d8786094f82dec19 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 09/17] fix[next]: address review of shared local dimensions in backends - iterator tracing and embedded shift name a sharing connectivity by offset_tag - 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 | 23 +++--- src/gt4py/next/iterator/tracing.py | 2 +- .../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 | 49 +++++++++++++ .../dace_tests/test_dace_utils.py | 31 ++++++++ .../unit_tests/test_neighbor_connectivity.py | 58 ++++++++++++++- 12 files changed, 242 insertions(+), 53 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 083c4d4359..5ffc813f8c 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1472,35 +1472,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: @@ -2160,14 +2170,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 73d7d00231..7062053223 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -582,7 +582,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 @@ -1012,9 +1012,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[ @@ -1467,7 +1468,7 @@ def _as_offset_tag( offset: runtime.Offset | type[common.NeighborConnectivity] | OffsetPart, ) -> OffsetPart: if isinstance(offset, common.ConnectivityMeta): - return common.local_dimension_of(offset).tag + return offset.offset_tag return offset.value if isinstance(offset, runtime.Offset) else offset @@ -1500,12 +1501,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) @@ -1560,9 +1563,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/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 094c85925f..ec8db888b1 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -154,7 +154,7 @@ def make_node(o): assert o.offset == 0 return im.cartesian_offset(o.domain_dim, o.codomain) if isinstance(o, common.ConnectivityMeta): - return OffsetLiteral(value=common.local_dimension_of(o).tag) + return OffsetLiteral(value=o.offset_tag) if callable(o): if o.__name__ == "": return lambdadef(o) 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 2ac65d0e58..3b44ecc4c2 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 @@ -769,7 +769,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 @@ -1441,7 +1443,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 da75325989..fb3099aa3d 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..d055f1f094 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 @@ -434,3 +434,52 @@ def test_sparse_shifted_stencil_reduce(program_processor): if validate: assert np.allclose(out.asnumpy(), ref) + + +class V2EShared(gtx.NeighborConnectivity[Vertex, Edge]): + """Shares `V2E`'s local dimension; bound to the table with its columns reversed.""" + + Local = V2EDim + + +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(V2EShared, 1)(in_edges)) + + +@fundef +def owner_times_sharer(in_edges): + return reduce(plus, 0)( + map_list(multiplies)(neighbors(V2EShared, 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={V2E.offset_tag: v2e_conn, V2EShared.offset_tag: 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..78f1045eb9 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 V2E, 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 = {V2E.offset_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 86bb41f1a3..a5d33b241c 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): ... 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'", ), ( """ @@ -249,6 +256,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] @@ -502,6 +518,46 @@ class Local(LocalDimensionIndex): ... assert not [w for w in recwarn if issubclass(w.category, DeprecationWarning)] +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 72039572a71da59b52d3d5dc18052cb019dd6d29 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 00:40:08 +0200 Subject: [PATCH 10/17] feat[next]!: class-keyed offset providers; remove FieldOffset Offset providers are keyed by NeighborConnectivity declarations. Every program entry point normalizes them to the tag-keyed form the IR and the backends use (as_tag_keyed_offset_provider), and tables are checked against their declarations once per compiled variant and on embedded calls (check_offset_provider). A bare string key is rejected as the removed FieldOffset spelling. FieldOffset is removed: unstructured connectivities are NeighborConnectivity declarations, Cartesian shifts are 'Dim + i', and as_offset takes the dimension to shift along. scripts/python/migrate_connectivities.py migrates user code. --- .../ADRs/next/0019-Connectivities.md | 4 + .../ADRs/next/0029-Connectivities_As_Types.md | 42 +- ...ectivities-as-types-implementation-plan.md | 1130 +++++++++++++++++ docs/user/next/QuickstartGuide.md | 8 +- .../exercises/2_divergence_exercise.ipynb | 2 +- .../2_divergence_exercise_solution.ipynb | 2 +- .../exercises/3_gradient_exercise.ipynb | 2 +- .../3_gradient_exercise_solution.ipynb | 2 +- .../workshop/exercises/4_curl_exercise.ipynb | 2 +- .../exercises/4_curl_exercise_solution.ipynb | 2 +- .../exercises/5_vector_laplace_exercise.ipynb | 8 +- .../5_vector_laplace_exercise_solution.ipynb | 8 +- .../8_diffusion_exercise_solution.ipynb | 2 +- docs/user/next/workshop/slides/slides_2.ipynb | 4 +- scripts/python/migrate_connectivities.py | 336 +++++ .../python/test_migrate_connectivities.py | 116 ++ src/gt4py/next/__init__.py | 2 - src/gt4py/next/common.py | 202 ++- src/gt4py/next/embedded/context.py | 6 +- src/gt4py/next/embedded/nd_array_field.py | 33 +- src/gt4py/next/ffront/decorator.py | 32 +- src/gt4py/next/ffront/experimental.py | 9 +- src/gt4py/next/ffront/fbuiltins.py | 93 -- .../ffront/foast_passes/type_deduction.py | 20 +- src/gt4py/next/ffront/foast_to_gtir.py | 6 +- src/gt4py/next/ffront/past_to_itir.py | 3 +- src/gt4py/next/ffront/transform_utils.py | 11 +- src/gt4py/next/iterator/embedded.py | 27 +- src/gt4py/next/iterator/runtime.py | 5 +- src/gt4py/next/otf/compiled_program.py | 1 + src/gt4py/next/otf/options.py | 10 +- .../codegens/gtfn/gtfn_module.py | 18 +- .../integration_tests/cases_utils.py | 53 +- .../ffront_tests/test_cartesian_shifts.py | 8 +- .../ffront_tests/test_compiled_program.py | 4 +- .../ffront_tests/test_concat_where.py | 12 +- .../ffront_tests/test_external_local_field.py | 16 +- .../ffront_tests/test_import_from_mod.py | 4 +- .../ffront_tests/test_named_collections.py | 4 +- .../test_neighbor_connectivity.py | 12 +- .../ffront_tests/test_reductions.py | 55 +- .../ffront_tests/test_staggered.py | 10 +- .../test_temporaries_with_sizes.py | 6 +- .../feature_tests/ffront_tests/test_tuples.py | 4 +- .../ffront_tests/test_type_conversion.py | 2 +- .../test_multiple_output_domains.py | 6 +- .../test_offset_dimensions_names.py | 118 +- .../embedded_tests/test_nd_array_field.py | 44 +- .../test_decorator_domain_deduction.py | 20 +- .../ffront_tests/test_foast_to_gtir.py | 20 +- .../ffront_tests/test_type_deduction.py | 23 +- .../runners_tests/dace_tests/test_dace.py | 2 +- .../dace_tests/test_dace_bindings.py | 2 +- .../dace_tests/test_dace_translation.py | 3 +- .../dace_tests/test_gtir_to_sdfg.py | 5 +- .../unit_tests/test_neighbor_connectivity.py | 94 +- typing_tests/test_next.yaml | 2 +- 57 files changed, 2140 insertions(+), 537 deletions(-) create mode 100644 docs/development/next/connectivities-as-types-implementation-plan.md create mode 100644 scripts/python/migrate_connectivities.py create mode 100644 scripts/tests/python/test_migrate_connectivities.py diff --git a/docs/development/ADRs/next/0019-Connectivities.md b/docs/development/ADRs/next/0019-Connectivities.md index 21827941bf..3b8296459f 100644 --- a/docs/development/ADRs/next/0019-Connectivities.md +++ b/docs/development/ADRs/next/0019-Connectivities.md @@ -9,6 +9,10 @@ tags: [] - **Created**: 2024-11-08 - **Updated**: 2026-05-27 +> The `FieldOffset` part of this record is superseded by +> [ADR 0029](0029-Connectivities_As_Types.md): connectivities are declared as +> `NeighborConnectivity` classes, and offset providers are keyed by them. + The representation of Connectivities (neighbor tables, `NeighborTableOffsetProvider`) and their identifier (offset tag, `FieldOffset`, etc.) was extended and modified based on the needs of different parts of the toolchain. Here we outline the ideas for consolidating the different closely-related concepts. ## History diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index e2b433fcd6..0bbc6bf82b 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-22 A neighbor connectivity is declared as a **class**, and its local dimension as a class **nested** in it: @@ -71,10 +71,8 @@ and skip values were never checked against the `FieldOffset` 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. + values whether or not an entry uses it. Programs run the check on the tables + they are given, see below. ### `NeighborConnectivity` is not a `Connectivity` @@ -139,6 +137,36 @@ works for both. The other frontend touch points treat the class like the 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. +### Offset providers are keyed by the declaration + +Users bind tables to declarations: + +```python +program(..., offset_provider={V2E: v2e_table, C2E: c2e_table}) +``` + +Every entry point of a program (`Program.__call__`, `FieldOperator.__call__`, +`compile`, `CompilationOptions.connectivities`, `embedded.context.update`, the +iterator `fendef`) normalizes such a provider to the form the IR uses: each +declaration is replaced by its `offset_tag`. Everything below the entry points — +lowering, the backends, compiled-program caching — therefore keeps seeing a +provider keyed by strings, which is also what hand-written IR uses. A string key +must be a tag, i.e. a qualified name; a bare name such as `"V2E"` is the removed +`FieldOffset` spelling and is rejected with a message pointing here. + +Tables are checked against their declarations (`check_offset_provider`) once per +compiled variant and on each embedded call, not on every compiled call: the +check builds the table's type, which is too slow for the call path. A tag that +names no declared connectivity, as in hand-written IR, is not checked. + +### `FieldOffset` is removed + +`FieldOffset` and its export are gone. An unstructured connectivity is a +`NeighborConnectivity`; a Cartesian shift is `Dim + i`, which the DSL already +had; and `as_offset` takes the dimension to shift along, `as_offset(KDim, k_offsets)`, instead of a Cartesian `FieldOffset`. `scripts/python/migrate_connectivities.py` +rewrites declarations and Cartesian offset uses, and reports the provider keys +and other sites it cannot rewrite from the source alone. + ## Consequences - An unstructured connectivity is spelled once. The provider key, the offset tag @@ -149,8 +177,8 @@ metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, - 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. +- `FieldOffset` is removed, and offset providers are keyed by declarations: a + breaking change for every unstructured program, eased by the migration script. ## Alternatives considered diff --git a/docs/development/next/connectivities-as-types-implementation-plan.md b/docs/development/next/connectivities-as-types-implementation-plan.md new file mode 100644 index 0000000000..296424e246 --- /dev/null +++ b/docs/development/next/connectivities-as-types-implementation-plan.md @@ -0,0 +1,1130 @@ +# Connectivities as types — implementation plan + +**Status**: **APPROVED** at revision 6 (adversarial review rounds 1–4; round 4 verdict APPROVED) +**Target**: an 8-PR stack on `main`, *alternative to* GridTools/gt4py#2844 +**Proposal**: `egparedes/connectivities-as-types` in GridTools/gt4py_knowledge (PR #32) +**Baseline tree**: `b3c53fa7e` (v1.2.2) + +## 0. Scope and relation to #2844 + +The proposal and #2844 agree on *what a dimension is* (a class, its indices its +instances) and disagree on *what identity a dimension has*. #2844 chose +`(tag, kind)` value equality with an interning registry; the proposal chose +nominal type identity with the tag being the qualified Python name. That single +disagreement propagates into five mechanisms, so the two cannot both land. + +This stack **re-cuts** #2844: it keeps the machinery independent of identity, +drops the machinery that exists only to support value identity, and then builds +the connectivity layer on top. **#2844 is closed, not merged** — that is what +makes this an alternative. Consequences: the new dimension ADR is **0028** (the +ADR directory on `main` ends at **0027**; 0028 exists only inside unmerged +#2844), and there is nothing to supersede. + +### Taken from #2844, unchanged in substance + +| Piece | Where in #2844 | +| -------------------------------------------------------------------------------------------------------------------------------- | ------------------------------------------ | +| `DimensionMeta` metaclass; `I + 1`, `I > 5`, `repr` living on it | `common.py` | +| `DimensionIndex` base: `__slots__ = ("value",)`, `kind` class keyword, `.dim` property | `common.py` | +| `type Dimension = type[DimensionIndex]` as a PEP 695 alias (so `Dimension("I")` raises rather than silently evaluating to `str`) | `common.py` | +| Metaclass `.value` property raising `AttributeError` that points at `.tag` | `common.py` | +| Deletion of `common.NamedIndex` (`.dim` / `.value` move onto the index instance) | `common.py` + ~40 call sites | +| Deletion of the dimension half of `mypy_plugin.py` (`_DimA`..`_AnyDim`); only the mixed-precision hooks remain | `type_system/mypy_plugin.py` | +| The mechanical migration of every declaration, incl. docs, workshop notebooks and `examples/` (which `test_examples` executes) | 131 files: 52 `src/`, 66 `tests/`, 13 docs | +| `xtyping.resolve_annotation` usage at `fbuiltins._type_conversion_helper` (already on `main` via #2841) | — | + +### Dropped from #2844 + +| Piece | Why | +| --------------------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | +| `_DIMENSION_REGISTRY` interning | identity is the type; nothing to intern | +| `copyreg.pickle(DimensionMeta, _reduce_dimension)` — the **blanket** registration on all dimensions | verified: a module-level dimension class pickles by reference with no help (`pickle.loads(pickle.dumps(KDim)) is KDim`). **But a narrow `copyreg` on `StaggeredMeta` is still required** — see §1.5 | +| The `DimensionMeta`-vs-`DimensionMeta` branch of `__eq__` / `__ne__` | becomes `is`. **The `IntegralScalar` overloads (`I == 5` → `Domain`) are kept**, and therefore so is an explicit `__hash__` — see §1.0 | +| `DimensionIndex.__eq__` comparing `type(self) == type(other)` | becomes `type(self) is type(other)` | +| `common.dimension(tag, kind)` factory | replaced by `common.resolve(tag)`, which imports | +| `fingerprinting.py` `DimensionMeta` deconstructor keyed on `(tag, kind)` | under type identity a dimension *is* fingerprinted by qualified name, so the generic `type` deconstructor is correct — **for the lenient variant only**. The STRICT variant rejects `Staggered[KDim]`, which is not importable under its qualified name. Both in-tree fingerprinters are lenient (`ffront/stages.py:62`, `iterator/ir.py:26`), and `eve_utils.content_hash` (`compiled_program.py:420`) is pickle-based and so needs §1.5's `copyreg`. Record the STRICT caveat in the ADR | +| ADR 0028 as drafted in #2844 | never lands; this stack writes its own 0028 | + +### Changed relative to #2844 + +| Piece | #2844 | This stack | +| --------------------------------- | ------------------------------------------------- | --------------------------------------------------------------------------------------------------------- | +| `tag` default | `cls.__name__`, settable in the class body | `f"{cls.__module__}.{cls.__qualname__}"`, a metaclass property; a class-body `tag = ...` is a `TypeError` | +| Rebuilding a dimension from a tag | `dimension(tag, kind)` (registry) | `resolve(tag)` (`import_module` + `qualname` walk), memoized | +| Declaration site requirement | none | module level, or unpicklable; `` heuristic in `__init_subclass__` | +| Backend name mangling | `tag` used directly | `codegen_name(tag)` + inverse, at ~19 enumerated sites in two name spaces (§1.3(b), (c)) | +| Staggered dimensions | `_Staggered` prefix through the interning factory | `Staggered[D]`, a real parametrized type — **required in PR 2**, not optional (§1.5) | + +### Superseded + +- **#2845** (`test[next]: adopt class-style dimension declarations`) is subsumed + by PR 2: because `dimension()` is not user-facing, the minimal + `I = gtx.dimension("I")` form does not exist and every declaration takes class + form immediately. #2845's pyright coverage is folded in. +- The `FieldOffset`-as-frontend-identifier part of **ADR 0019**. +- **ADR 0026**'s `_Staggered` name prefix (PR 2). + +## 1. Design questions closed before implementation + +Everything in this section was verified by running it, not by reading. Probe +files are named; they become committed test material in the PR that needs them. + +### 1.0 Metaclass mechanics that are easy to get wrong + +**`__hash__` must be declared explicitly.** Python sets `__hash__ = None` on any +class body that defines `__eq__` without `__hash__` — metaclasses included. Since +the `I == 5` → `Domain` overload keeps `__eq__` on `DimensionMeta`, dropping +#2844's `__hash__` makes every dimension class *unhashable*: + +``` +>>> class M(type): +... def __eq__(cls, o): return True +>>> M.__hash__ is None +True +>>> class C(metaclass=M): pass +>>> hash(C) +TypeError: unhashable type: 'M' +``` + +That would break `domain({I: 2})` (`common.py:672-690`), +`Counter[common.Dimension]` (`embedded/nd_array_field.py:314`), +`dict[Dimension, SymbolicRange]` (`iterator/ir_utils/domain_utils.py:136,152`), +`seen: dict[Dimension, Dimension]` (`common.py:1351`), and eve's validator +memoization on annotation objects (`eve/type_validation.py:599`) — so +`ts.DimensionType` would fail at *import*. **Fix: `__hash__ = type.__hash__` +explicitly on `DimensionMeta`,** and likewise on `ConnectivityMeta` if it ever +defines `__eq__`. + +**A metaclass `__getitem__` shadows `__class_getitem__`.** `ConnectivityMeta` +needs `__getitem__` for `V2E[1]` (the single-neighbor shift handle that +`FieldOffset.__getitem__` provides today), but metaclass lookup takes precedence +over `Generic.__class_getitem__`, so a naive implementation makes +`NeighborConnectivity[V, E]` in a bases list fail with +`TypeError: tuple expected at most 1 argument, got 3`. + +**Fix, verified clean under `mypy --strict` and pyright 1.1.414 on Python 3.12** +(`/tmp/probe_meta_getitem3.py`): dispatch on the argument type, delegating +non-`int` subscription back to `cls.__class_getitem__`: + +```python +class ConnectivityMeta(type): + __hash__ = type.__hash__ + + @overload + def __getitem__(cls, item: int) -> Connectivity: ... + @overload + def __getitem__(cls, item: Any) -> Any: ... + def __getitem__(cls, item: Any) -> Any: + # `numbers.Integral`, not `int`: `V2E[np.int32(1)]` must not fall through + # to the type-parameter branch (it raises `TypeError: V2E is not a + # generic class` there). `bool` is excluded so `V2E[True]` is an error + # rather than silently neighbor 1. + if isinstance(item, numbers.Integral) and not isinstance(item, bool): + return _bound_single_neighbor(cls, int(item)) + # type-parameter subscription, e.g. `NeighborConnectivity[V, E]` + return cast(Any, cls).__class_getitem__(item) +``` + +`cast(Any, cls)`, not `super()` — `__class_getitem__` is on the class, not on the +metaclass MRO; `super().__class_getitem__` raises `AttributeError`. With the cast +both checkers report zero errors and all uses work at runtime +(`NC[V, E]`, `class V2E(NC[V, E])`, `V2E[1]`, and `V2E.Local` as an annotation). +pyright accepts `NC[V, E]` in a **bases list**; in a *value* position it types it +`Any`, which is why the overloads above matter — without them `V2E[1]` is also +`Any` and the shift handle is untyped. + +### 1.1 `NeighborConnectivity` is **not** a `Connectivity` (proposal Open Q6) + +`common.Connectivity` is `Field[DimsT, IntegralScalar]` — a **data** protocol +(`common.py:990`; `ndarray`, `asnumpy`, `domain` are all on it). A declaration +class holds no data. + +**Resolution.** Two distinct things, distinct hierarchies: + +- `NeighborConnectivity` — a **declaration**. Not a `Connectivity`. It produces a + `NeighborConnectivityType` via `__gt_type__()`, is the provider key, and is the + handle written in DSL code (`a(V2E)`). +- `NeighborTable` / `NdArrayConnectivityField` — the **data**, unchanged, still + `Connectivity` implementations. + +This is the shape `FieldOffset` already has: it is *not* a `Connectivity` either, +and `premap` special-cases it at `nd_array_field.py:317-320`. So `a(V2E)` +continues to work by widening the same union — `Field.premap` and +`Field.__call__` are typed `Connectivity | fbuiltins.FieldOffset` +(`common.py:785, 791-794`) and become `Connectivity | type[NeighborConnectivity]` +in PR 4. `V2E` has **no instances**: `ConnectivityMeta.__call__` raises +`TypeError("… is a connectivity declaration and cannot be instantiated; bind a table through the offset provider")`. + +Consequence: the proposal's sketch line +`class NeighborConnectivity(Connectivity[MultiDimensionIndex[Origin, Local], Codomain], ...)` +is **wrong and dropped**. `MultiDimensionIndex` remains the *domain index type* of +the `NeighborTable` (PR 8). **The knowledge-repo note needs this correction.** + +### 1.2 How `Local` reaches the base (proposal Open Q2) + +`requires-python = '>=3.12'`, so a PEP 696 default type parameter (3.13) is not +available. Resolution: **metaclass discovery**, base carrying a `ClassVar` +annotation, subclass declaring the nested class explicitly: + +```python +class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + Local: ClassVar[type[LocalDimensionIndex]] # annotation only, never assigned + + +class V2E(NeighborConnectivity[V, E], max_neighbors=6): + class Local(LocalDimensionIndex): ... # explicit, required +``` + +Verified under `mypy --strict --python-version 3.12` and `pyright --pythonversion 3.12` +(`/tmp/probe_local.py`): + +| Variant | base declares | mypy | pyright | +| ------- | ------------------------------------------------ | ----- | ---------------------------------------- | +| 1 | `Local: ClassVar[type[LocalDimensionIndex]]` | clean | clean | +| 2 | nothing | clean | clean | +| 3 | a real nested `class Local(LocalDimensionIndex)` | clean | **`reportIncompatibleVariableOverride`** | + +In all three the intended negative case (`Field[V, A.Local]` vs +`Field[V, B.Local]`) is correctly an error. Variant 3 is rejected. + +**Stated precisely — what variant 1 does and does not buy.** It does *not* make +`conn.Local` usable as a **type annotation** when `conn` is a generic +`type[NeighborConnectivity]`: both checkers reject that (mypy `name-defined`, +pyright `reportInvalidTypeForm`), and `T.Local` on a `TypeVar` is rejected too. +What variant 1 buys over variant 2 is only **value-level** access — +`reveal_type(conn.Local)` is `type[LocalDimensionIndex]` instead of an attribute +error — which is what library code in `common`, the backends and +`type_synthesizer` actually needs. Variant 1 is chosen for that, not for generic +annotations. Generic library code that must *name* a local dimension in a +signature uses `type[LocalDimensionIndex]`. + +This extends the proposal's probe P2: a **generated** `Local` is unusable as an +annotation, but a base `ClassVar` *annotation* plus an explicitly declared nested +class is fine. + +### 1.3 The IR keeps string tags; `resolve()` and `codegen_name()` are both required + +`AxisLiteral.value: str` stays (making it carry the class is a separate IR +change, deferred past this stack). It now holds the **qualified** tag, and that +has two consequences the first draft of this plan underestimated. + +**(a) `resolve(tag)` at every rebuild site**, memoized — `inference.py:464` calls +it once per `AxisLiteral` on the type-inference hot path: + +| Site | Purpose | +| ------------------------------------------- | -------------------------------------------------------------------------------------------------------------------------------- | +| `iterator/ir_utils/domain_utils.py` | `AxisLiteral` → `Dimension` | +| `iterator/ir_utils/misc.py` | `AxisLiteral` → `Dimension` | +| `iterator/type_system/inference.py:464` | `AxisLiteral` → `ts.DimensionType` | +| `codegens/gtfn/itir_to_gtfn_ir.py` (×2) | staggered-name sniffing → replaced in PR 2 by `Staggered[D]` | +| `dace/lowering/gtir_to_sdfg_lambda.py:1155` | synthesizes the local dim from the *offset* tag: `Dimension(offset, LOCAL)`. **Must be fixed in PR 2, not deferred** — see below | +| `dace/sdfg_args.py` | axis name → `Dimension` | +| `runners/roundtrip.py` | emits `gtx.Dimension(...)` as *source text* → becomes an import | +| ~~`common.flip_staggered` (×2)~~ | **not** a `resolve()` site: `Staggered[D]` replaces it with an interning subscript, §1.5 | + +`resolve` on a nested qualname was verified to work and round-trip +(`resolve("mymod.V2E.Local") is mymod.V2E.Local`), which matters because PR 4 keys +the provider on `V2E.Local.tag`. **One hazard to settle in PR 2**: a purely dotted +tag does not record *where* the module path ends and the qualname begins, so +`resolve` must try the longest importable prefix and walk the rest — O(depth) +import attempts, and in principle ambiguous if a module path and a class-attribute +chain collide. `pickle` avoids this by storing module and qualname *separately*. +Options: keep the pure dotted form (what the proposal asks for, ambiguity +tolerated and memoized away) or use an explicit separator such as +`"module:qualname"`. **Recommendation: keep the dotted form** — it is what makes +the tag "also a valid tag string for the IR", the collision requires a module and +an attribute chain to have the same spelling, and `resolve` can prefer the +*longest* importable prefix so a real module always wins. Record the residual in +the ADR. + +**(b) `codegen_name(tag)` — dots are illegal in every generated identifier.** +`eve`'s `SymbolName`/`SymbolRef` are constrained by +`_SYMBOL_NAME_RE = ^[a-zA-Z_]\w*$` (`eve/concepts.py:23,26,32`), so a qualified +tag reaching `Sym(id=...)` is a *validation error*, not a cosmetic problem. The +first draft mentioned mangling only in the abstract and put the roundtrip change +in a later PR; both were wrong. All of these are **PR 2**: + +| Site | What breaks without mangling | +| -------------------------------------------------------- | ----------------------------------------------------------------------------------------------------------- | +| `codegens/gtfn/itir_to_gtfn_ir.py:170-195` | `TagDefinition(name=Sym(id=dim.value))` → `SymbolName` validation error | +| `codegens/gtfn/gtfn_module.py:97, 130-136` | `generated::{dim.value}_t`, plus `name.lower()` | +| `otf/binding/nanobind.py:197, 211` | C++ identifiers | +| `dace/lowering/gtir_to_sdfg_utils.py` `get_map_variable` | `i_{dim.value}_gtx_{kind}` → invalid DaCe symbol | +| `dace/sdfg_args.py:80` `_field_symbol` | invalid DaCe symbol | +| `dace/lowering/gtir_python_codegen.py:137-138` | `visit_AxisLiteral` returns the raw value | +| `runners/roundtrip.py:64, 177` | `AxisLiteral = as_fmt("{value}")`, and `{o.value} = gtx.Dimension(...)` emits `a.b.I = ...` → `SyntaxError` | + +**(d) The mangling scheme, corrected.** Earlier drafts said "injective (escape +existing `__` before replacing `.`)", i.e. `_ -> __` then `. -> _`. **That is not +injective**: `.` becomes a single `_`, so `".."` and `"_"` both map to `"__"`. +Exhaustively tested over the alphabet `{a, ., _}` up to length 6 +(`/tmp/probe_mangle.py`): **686 collisions in 1092 inputs.** Since a generated +identifier may only contain `[A-Za-z0-9_]`, `_` is the only available separator +and a *prefix escape* is required: + +```python +def codegen_name(tag: Tag) -> str: + return tag.replace("_", "_u").replace(".", "_d") + + +def from_codegen_name(name: str) -> Tag: + return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) +``` + +Every `_` in the output is the first character of a two-character escape, so +decoding is unambiguous. Verified exhaustively over `{a, ., _, u, d}` up to +length 6 — **19530 inputs, 0 collisions, 0 round-trip failures**, every output a +valid identifier, including the adversarial `"_u"`, `"_d"` and `"a_ud.b"` +(`/tmp/probe_mangle2.py`). Cost: names grow (`mod.V2E.Local` → +`mod_dV2E_dLocal`), which is what gtfn's existing `TagDefinition.alias` mechanism +is for. + +**An inverse is needed too**, wherever generated names are parsed *back* into +dimensions: `dace/sdfg_args.py:25, 60-72` matches `gt_conn_(\S+)` and feeds the +result to `has_offset`. `codegen_name` must therefore be injective *and* have a +`from_codegen_name` partner (escape `__` → `____` before `.` → `__`). + +**A site that cannot be deferred: `gtir_to_sdfg_lambda.py:1155`.** It builds +`gtx_common.Dimension(offset, DimensionKind.LOCAL)` — a local dimension +synthesized from the **offset** tag, which in PR 2 is still a bare provider key +(`"V2E"`) that `resolve()` cannot import. Every DaCe unstructured shift passes +through it, so PR 2 is red on DaCe unless it is fixed there. The fix is local and +available: `conn_type` is already in scope (`:1134-1152`) and `:1135` already +asserts `conn_type.domain[1].kind == LOCAL`, so the line becomes +`offset_type = conn_type.domain[1]` (equivalently `conn_type.neighbor_dim`). +It is *necessary but not sufficient* for PR 1's `shift × tag≠localdim` DaCe cell: +that cell fails earlier, at `gtir_to_sdfg.py:842` +(`neighbor_table_types[dim.value]`, i.e. A4 on the connectivity *argument's* local +dim), before `:1155` is reached — and after the `:1155` fix, `:1371`/`:1455` would +reference `gt_conn_` while `:1104`/`:722` declare `gt_conn_`. So +**the DaCe shift cell stays in the skip matrix until PR 4**, where the +single-string choice makes both agree. (An earlier draft said PR 2; that would +leave PR 2 red on that cell.) + +**(c) The *offset* key is a second dotted name space, and it is mangled in PR 4, +not PR 2.** §1.3(b) covers *dimension* names only. When PR 4 makes the provider +key `cls.tag`, the **offset** string that flows through the IR +(`OffsetLiteral.value`, the provider key, `o` in the gtfn/DaCe connectivity +plumbing) becomes dotted too, and a different set of sites turns *it* into an +identifier. These are all **PR 4**: + +| Site | What breaks | +| ------------------------------------------------------------------------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `codegens/gtfn/itir_to_gtfn_ir.py:184` | `TagDefinition(name=Sym(id=offset_name))` → `SymbolName` regex | +| `codegens/gtfn/itir_to_gtfn_ir.py:490` | `SymRef(id=o)` for each connectivity → `SymbolRef` regex | +| `codegens/gtfn/codegen.py:147-148` | `visit_OffsetLiteral` emits `node.value` raw into C++ | +| `codegens/gtfn/gtfn_module.py:118, 132, 136` | `GENERATED_CONNECTIVITY_PARAM_PREFIX + name.lower()`, `generated::{name}_t` | +| `dace/sdfg_args.py:56` | `connectivity_identifier(name)` → `gt_conn_a.b.V2E`, an invalid SDFG array name | +| `dace/sdfg_args.py:60`, `dace/workflow/bindings.py:200, 286` | `is_connectivity_identifier` / `_parse_gt_connectivities` — the **inverse** direction, so `from_codegen_name` has *several* live consumers, not one | +| `dace/workflow/translation.py:61`, `dace/sdfg_callable.py:103`, `dace/program.py:156` | `connectivity_identifier(offset)` again, on the argument-binding path | +| `dace/lowering/gtir_to_sdfg_lambda.py:1104, 1371, 1455, 1727`, `gtir_to_sdfg.py:722` | the same identifier, consumed in the lowering | +| `runners/roundtrip.py:63` | `OffsetLiteral = as_fmt("{value}")` — emits the offset tag *raw as Python source*, into the program **body**; mangling `:176` alone still leaves `NameError: name 'tests' is not defined` | +| `dace/sdfg_args.py:83-84` | `_field_symbol`: `assert m[1] in offset_provider_type` — a *second* `from_codegen_name` consumer besides `:70` | +| `dace/lowering/gtir_to_sdfg_lambda.py:1892` | `visit_OffsetLiteral` → `SymbolExpr(node.value, INDEX_DTYPE)`, i.e. a dotted string used as a DaCe symbolic expression | +| `runners/roundtrip.py:152, 176` | collects offset-literal strings, then `f'{o} = offset("{o}")'` → `a.b.V2E = offset(...)` → `SyntaxError` | + +So `codegen_name` / `from_codegen_name` are introduced in PR 2 for dimensions and +**applied again in PR 4 for offsets**, at ~16 further sites. Two earlier claims +were wrong: that `from_codegen_name`'s only live consumer is in PR 2, and that the +DaCe surface is confined to `sdfg_args.py` and the lowering — the +argument-binding path (`workflow/translation.py`, `workflow/bindings.py`, +`sdfg_callable.py`, `program.py`) carries it too, in both directions. + +Because of (b) and (c), the review shortcut "diff PR 2 against #2844, the delta is +only identity" is **false**: #2844 needed none of this. Reviewers should expect a real +mangling layer on top of the identity delta. + +**`AxisLiteral.kind` becomes redundant** (the class carries it) — the `TODO` at +`iterator/ir.py:93`. Kept in PR 2, removed in PR 7, to keep PR 2's IR-expectation +churn to the `value` strings only. + +### 1.4 `LocalDimensionIndex` subclasses `DimensionIndex`; `DimensionBaseIndex` is dropped + +The proposal lists `DimensionBaseIndex` as a separate root with `DimensionIndex` +and `LocalDimensionIndex` as siblings. That does not survive contact with the +tree: `Dimension` is `type[DimensionIndex]`, eve validates a `type[X]` +annotation by `issubclass` (verified: a subclass passes, the base and an +unrelated class are both rejected), so sibling local dimensions would force +widening to `type[DimensionBaseIndex]` at `ts.DimensionType.dim`, +`ts.FieldType.dims`, `ConnectivityType.domain`, `Domain.__init__` and the `DimT` +/ `DimT_co` bounds — and would then accept local dimensions everywhere a primary +one is meant, which is the same looseness with extra ceremony. + +**Resolution**: `class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL)`. +`DimensionBaseIndex` is not introduced at all — one concept fewer, which is the +proposal's own stated goal. Where primary-only is required the check is +`dim.kind is not DimensionKind.LOCAL`, exactly as today. This also removes +#2844's deferral note ("a `DimensionBase` root, deferred until the requirements +of non-user-declarable dimensions are known") as a thing that needs resolving. + +**Verified**: all 38 sites in `src/` that discriminate a local dimension do so by +a **runtime `kind` check**, not by a static type distinction +(`transform_utils.py:65`, `type_deduction.py:460, 774`, +`custom_layout_allocators.py:171`, `past_to_itir.py:409`, `common.py:1168, 1336`, +`nd_array_field.py:972, 976`, `gtfn_module.py:91`, `embedded.py:922`, +`gtir_to_sdfg_types.py:76`, …). The tree already treats local dimensions as +`Dimension`s everywhere — `ConnectivityType.domain: tuple[Dimension, ...]` +includes the local one — so subclassing loses nothing it currently relies on, and +`Dims` (`tuple[Unpack[ShapeTs]]`, `common.py:57`) puts no bound on its members +either. + +**What subclassing does cost**, and the mitigation: every `DimensionIndex` +*bound* now statically admits a local dimension — +`NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]`, +`Staggered[D: DimensionIndex]` and `MultiDimensionIndex[D: DimensionIndex, *Ls]` +would all accept `V2E.Local` as their primary parameter. Each therefore gets a +runtime `kind is not DimensionKind.LOCAL` check in `__init_subclass__` / +`__class_getitem__`, and `LocalDimensionIndex.__init_subclass__` rejects an +explicit `kind=` other than `LOCAL`. This is the same runtime-check discipline +the tree already uses; the static gap is the price of the concept removed. + +**Deviation from the proposal; needs feeding back to the note.** + +### 1.5 `Staggered[D]` is required in PR 2, not PR 7 + +`flip_staggered` builds `Dimension(f"_Staggered{name}")` from a string +(`common.py:1452-1457`) and `is_staggered` tests `dim.value.startswith(prefix)` +(`:1447-1449`). #2844 routes both through the interning factory. With the +registry gone there is **no importable `_Staggered` type**, and a +dynamically created class would get the tag +`gt4py.next.common._Staggered`, so `is_staggered` is false and +`as_non_staggered` cannot recover the base dimension's module. Live dependents: +`test_staggered.py` (233 lines), `cases_utils.py:161` +(`KHalfDim = flip_staggered(KDim)`), gtfn `_add_staggered_aliases` +(`itir_to_gtfn_ir.py:203-215`), DaCe `get_map_variable` +(`gtir_to_sdfg_utils.py:52`), `type_synthesizer`, `test_common.py`, +`test_domain_utils.py`. + +So PR 2 is **not green** without `Staggered[D]`. It is Cartesian-only and does +not depend on the connectivity layer, so it moves into PR 2. + +**The obvious mechanism does not work.** A PEP 695 generic +`class Staggered[D: DimensionIndex](DimensionIndex)` makes `Staggered[KDim]` a +`typing._GenericAlias`, **not a class** (verified, `/tmp/probe_staggered.py`): + +``` +type(Staggered[KDim]) -> +isinstance(Staggered[KDim], type)-> False +issubclass(Staggered[KDim], ...) -> TypeError: issubclass() arg 1 must be a class +Staggered[KDim].tag -> '__main__.Staggered' # KDim is gone +``` + +So it fails eve's `type[DimensionIndex]` validation and its tag cannot name the +base dimension — it is not a `Dimension` at all. + +**The mechanism that does work** (verified, `/tmp/probe_staggered3.py`: runs +correctly and is **0 errors under both `mypy --strict` and pyright 1.1.414** on +3.12) is a metaclass `__getitem__` that *builds and interns a real class*, paired +with a `TYPE_CHECKING` declaration so checkers still see an ordinary generic: + +```python +class StaggeredMeta(DimensionMeta): + def __getitem__(cls, base: Dimension) -> Dimension: + if base not in _staggered_cache: + _staggered_cache[base] = StaggeredMeta( + f"Staggered[{base.__name__}]", + (cls,), # NOT (cls, base) -- see below + { + "_tag": f"{cls.__module__}.{cls.__qualname__}[{base.tag}]", + "kind": base.kind, + "base": base, + "__slots__": (), + }, + ) + return _staggered_cache[base] + + +if TYPE_CHECKING: + + class Staggered[D: DimensionIndex](DimensionIndex): + base: ClassVar[Dimension] +else: + + class Staggered(DimensionIndex, metaclass=StaggeredMeta): + __slots__ = () + base: ClassVar[Dimension] +``` + +Verified properties of `Staggered[KDim]`: it *is* a class; +`tag == "gt4py.next.common.Staggered[]"`; `kind` is +inherited from the base; `issubclass(_, DimensionIndex)` and +`issubclass(_, Staggered)` hold; it is instantiable as an index; and +`Staggered[KDim] is Staggered[KDim]`, so identity is stable. `Staggered[KDim]` in +an annotation and inside `Field[Dims[Staggered[KDim]], float]` are both accepted +by both checkers. + +- **Bases are `(cls,)`, not `(cls, base)`.** Inheriting from the base dimension + would make `issubclass(Staggered[KDim], KDim)` true, i.e. `KHalfDim` would be + accepted everywhere `KDim` is required. It is a *different* dimension; only + `kind` is inherited, copied explicitly into the namespace. +- `is_staggered(dim)` becomes **`"base" in dim.__dict__`**, not + `issubclass(dim, Staggered)`, and `as_non_staggered(dim)` becomes `dim.base`. + Two runtime facts force this: `issubclass(Staggered, Staggered)` is true for + the bare base, which has no `base`; and `Staggered[KDim]` is *subclassable* + (`class KHalf2(Staggered[KDim])` yields a second, un-interned staggered-K type + with tag `.KHalf2`). `Staggered.__init_subclass__` therefore rejects + any subclass the metaclass did not create, so the interned form is the only + one. Still structural — no string sniffing. +- **The guards were verified, including the escape routes** + (`/tmp/probe_staggered_guards.py`). All four are blocked: + `class KHalf2(Staggered[KDim])`, `class X(Staggered)`, a direct + `StaggeredMeta("Y", (Staggered,), {})`, and the double subscript + `Staggered[KDim][KDim]`. The `copyreg` fallback round-trips the bare + `Staggered` by reference, the parametrized class with identity preserved, and + instances. Implementation note: gate `__init_subclass__` on a **namespace + marker** the metaclass sets (`"_tag" in cls.__dict__`), not on a module-level + "currently building" flag — the flag works but is not thread-safe, and + compilation runs in worker processes and threads. Three further honest limits: + the guard defends against **accidental** subclassing only — a deliberate + `StaggeredMeta("Forged", (Staggered,), {...marker})` or `types.new_class` can + still forge a same-`tag`, non-identical type (as it can for any class); + `Staggered[Staggered[KDim]]` must be rejected explicitly by testing + `"base" in base.__dict__` in `__getitem__`, or it nests and pickles happily; and + a hand-built `copyreg` payload such as `(_make_staggered, (int,))` should raise + a `TypeError` naming the offending base rather than an `AttributeError`. +- `resolve` gains the `[]` grammar: it parses the brackets and + evaluates `Staggered[resolve(inner)]`, which hits the same intern cache, so a + staggered dimension round-trips through the IR to the *same* class object. +- **A narrow `copyreg` is required after all.** `Staggered[KDim]`'s + `__qualname__` is `Staggered[KDim]`, which `pickle.save_global` cannot look up: + `PicklingError: Can't pickle : attribute lookup Staggered[KDim] on … failed` (verified, `/tmp/probe_staggered_pickle.py`). A + `copyreg.pickle(StaggeredMeta, lambda cls: (_make_staggered, (cls.base,)))` + fixes it *and preserves identity*, because the reconstructor goes back through + the intern cache. **But it must guard the bare base**: `type(Staggered) is StaggeredMeta` too, so a reducer that unconditionally reads `cls.base` fails on + `Staggered` itself with `AttributeError: type object 'Staggered' has no attribute 'base'` (verified — an earlier draft of this section claimed the + registration "captures only parametrized dimensions", which is false). The + reducer therefore falls back to by-reference pickling when + `"base" not in cls.__dict__`. It never captures a plain dimension + (`type(KDim) is DimensionMeta`). This is materially narrower than + #2844's blanket registration on `DimensionMeta` — a parametrized type needs a + reconstructor for the same reason `typing` aliases do — but §0's "`copyreg` + dropped" row is only true of the blanket form, and the ADR must say so. +- **Two honest costs.** (i) `_staggered_cache` is a cache, and the proposal's + headline is that the *name-keyed* registry goes away. The difference is real but must be + stated: it is keyed by a *dimension class*, is internal, and is memoization of + a type constructor (as `typing`'s own subscription cache is), not interning of + user-authored name strings — nothing resolves a user string through it. + (ii) the `TYPE_CHECKING` split means the static and runtime definitions can + drift; a unit test must assert the runtime facts the static form does not + express (real class, `issubclass` against `Staggered` but *not* against the + base, tag shape, interning). +- Supersedes ADR 0026, recorded in the PR-2 ADR. + +## 2. The PR stack + +Branches follow the repo's stacked convention, `connectivities-as-types--`, +each based on its predecessor, all targeting `main`. PR titles are Conventional +Commits (squash-merge lands the title). + +______________________________________________________________________ + +### PR 1 — `fix[next]: lower unstructured shifts with the offset's own tag` + +**Independent of the rest of the stack; lands first, on its own merit.** + +`foast_to_gtir._visit_shift` emits the **Python variable name** as the IR shift +tag (`foast_to_gtir.py:305` `offset_name.id`, `:331` `str(offset_name)`), because +`ts.OffsetType` does not carry the tag. Embedded execution keys on +`FieldOffset.value`. So the same program needs a *different* provider key +depending on the backend — confirmed by running it on v1.2.2: + +``` +MyOff = FieldOffset("TAGNAME", ...) +embedded: {"TAGNAME": conn} OK ; {"MyOff": conn} -> KeyError 'TAGNAME' +roundtrip: {"MyOff": conn} OK ; {"TAGNAME": conn} -> KeyError 'MyOff' +``` + +**Change** + +- `ts.OffsetType` gains **`tag: Optional[Tag] = None`** — *not* a required field. + `type_deduction.py:709` builds `ts.OffsetType(source=conn.codomain, target=(conn.domain_dim,))` from `IDim + 1`, a `CartesianConnectivity` that has + no tag at all; making `tag` required breaks it. +- `FieldOffset.__gt_type__` fills it (`fbuiltins.py:485`). +- `type_deduction.py:464`, which rebuilds an `OffsetType` when `Off[1]` drops the + local dimension, must **propagate** the tag. +- `foast_to_gtir._visit_shift`: the `Subscript` branch and the bare `Name` branch + use `arg.type.tag`, asserting non-`None` (both are unstructured paths, where a + tag always exists). + +**Tests.** `tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py` +today covers exactly `a(Off[1])` on `GTFN_CPU`. Extend to +{shift, `neighbor_sum`} × {embedded, roundtrip, gtfn, dace} × {tag≠varname, +tag≠local-dim-name}. + +**The matrix is not uniform, and a blanket `xfail` will not do.** `xfail_strict = true` +(`pyproject.toml:323`), and measured behaviour on v1.2.2 is: + +| case | embedded | roundtrip | gtfn | dace | +| ---------------------------- | -------- | --------- | -------- | ------------------------------------------------------------------------------------------ | +| shift, tag≠varname | pass | pass | pass | pass | +| shift, tag≠localdim | pass | pass | pass | **fail** `KeyError` (`gtir_to_sdfg_lambda.py:1155` synthesizes the local dim from the tag) | +| `neighbor_sum`, tag≠localdim | **fail** | **pass** | **fail** | **fail** | + +So the first draft's acceptance criterion ("shift cells pass on all four +backends") is unreachable before the backend work, and a strict blanket `xfail` +would XPASS on roundtrip. **Fix**: add a per-backend skip matrix entry in +`tests/next_tests/definitions.py` (a new `USES_*` marker) covering exactly the +failing cells, roundtrip excluded. **They are removed in two steps**: the +`shift × tag≠localdim` DaCe cell and the three `neighbor_sum × tag≠localdim` cells +all in **PR 4**, where the single-string choice makes A3/A4 vacuous — *not* in +PR 5, and *not* the DaCe cell in PR 2 (the `:1155` fix there is necessary but not +sufficient; see §1.3(a)). The gtfn `neighbor_sum` +failure is now confirmed **by running it**; the proposal had it only "by +reading". + +**No CHANGELOG entry.** Verified against the history: `CHANGELOG.md` is touched +*only* by release PRs (`git log -- CHANGELOG.md` is release commits exclusively, +and nothing between `b3c53fa7e` and `upstream/main` touches it). The behaviour +change — which key a compiled backend requires when tag ≠ variable name — belongs +in the PR description, and reaches the changelog when the release PR is cut. Two +earlier drafts of this plan said otherwise, including for PR 6's breaking change. + +**ICON4Py is unaffected by PR 1**: all 16 `FieldOffset` variable names equal +their tags (`model/common/src/icon4py/model/common/dimension.py:33-48`). + +**Acceptance**: `nox -s test_next` green; every cell in the matrix either passes +or is covered by the documented skip matrix. + +______________________________________________________________________ + +### PR 2 — `feat[next]: a concrete Dimension is a class, identified by its qualified name` + +The #2844 core with the identity divergences of §0, **plus** the mangling layer +of §1.3(b) and `Staggered[D]` of §1.5 — both of which #2844 did not need and +without which this PR cannot be green. Large and largely mechanical. + +**`src/gt4py/next/common.py`** + +```python +class DimensionMeta(type): + kind: DimensionKind + __hash__ = type.__hash__ # §1.0 — mandatory, not optional + + @property + def tag(cls) -> Tag: ... # f"{cls.__module__}.{cls.__qualname__}" + + # operators as in #2844; __eq__/__ne__ keep only the IntegralScalar overload + # (I == 5 -> Domain); the dim-vs-dim branch is `is`. + + +class DimensionIndex(metaclass=DimensionMeta): + __slots__ = ("value",) + kind: ClassVar[DimensionKind] = DimensionKind.HORIZONTAL + + def __init_subclass__(cls, /, kind=None, **kw): ... + + +# Staggered: an interning metaclass + TYPE_CHECKING split, NOT a PEP 695 +# generic -- see §1.5, where the generic form is shown to be unworkable. + + +def resolve(tag: Tag) -> Dimension: ... # memoized; [] grammar +def codegen_name(tag: Tag) -> str: ... # "_" -> "_u", "." -> "_d" (§1.3(d)) +def from_codegen_name(name: str) -> Tag: ... # the inverse, `_([ud])` -> `_` / `.` + + +type Dimension = type[DimensionIndex] +``` + +- `tag` is a metaclass **property**, so it cannot drift from the type. This makes + a class-body `tag = "C2E"` a **silent no-op** (verified: `C2EDim.tag` stays + `"__main__.C2EDim"` even with `tag = "C2E"` in the body) — and that pattern is + exactly what ICON4Py and #2845 use to rename. `__init_subclass__` therefore + **raises** on `"tag" in cls.__dict__`, naming the class and pointing at the + rename path. +- `__init_subclass__` also rejects `"" in cls.__qualname__`. Neither + necessary nor sufficient (`type("Dyn", ...)` in a function passes; a `del`'d + class passes) — the authoritative check stays pickle's own `save_global`. +- `resolve` raises a `ValueError` naming the tag and the failing import, per + CODING_GUIDELINES. + +**Removals**: `NamedIndex`; `_DimA`..`_AnyDim` and the dimension half of +`mypy_plugin.py`; `_DIMENSION_REGISTRY`; `copyreg`; the `fingerprinting.py` +deconstructor; `_STAGGERED_PREFIX` and its string sniffing. + +**Migration**. Every `Dimension("X")` becomes `class X(DimensionIndex): ...` at +module level. Verified counts: 333 `Dimension("` declarations in `tests/`, of +which **133 are function-local across 15 files** and must move to module level; +a dimension *named* `"I"` is declared 46 times across **17** files (the first +draft said 45 files — that was the proposal's *`IDim` file* count, a different +number). Docs, workshop notebooks and `examples/` are included because +`test_examples` executes them; notebook *code* cells only, stored outputs +untouched (they hold recorded tracebacks that must keep naming the symbols that +produced them). + +**IR expectation churn**: 36 `AxisLiteral` and 37 `OffsetLiteral` occurrences in +`tests/`, most already computed from `dim.value`. `test_pretty_roundtrip.py` and +the gtfn/DaCe snapshot tests hold the hardcoded names. + +**Do the sweep with a codemod script, not agent fan-out.** A previous attempt at +agent fan-out on a large mechanical rewrite in this repo died mid-file on the +rate limit and left the tree inconsistent; a script did all 57 files uniformly. + +**ADR 0028** (the directory ends at 0027): nominal identity; the module-level +declaration requirement; `Staggered[D]` superseding ADR 0026; that cache +fingerprints now shift when a declaration moves module (a consequence for ADR +0023, not a reversal); that `resolve()` imports modules named in the IR, which is +the same trust level as `pickle` loading a class by reference. + +**Documented limitation**: interactive `__main__` (REPL, notebooks, `python -c`) +cannot be resolved. `spawn` compile workers re-execute the main *script* as +`__mp_main__`, so file-based `__main__` resolves provided the script has the +`if __name__ == "__main__":` guard the pool already requires. + +**ICON4Py migration script is a PR-2 deliverable, not PR 6.** All 15 local +dimensions and `KDim`/`EdgeDim`/`CellDim`/`VertexDim` have variable name ≠ tag +(`EdgeDim = Dimension("Edge")`), so PR 2 changes every generated symbol and every +cache key downstream. + +**Acceptance**: `nox -s test_next` on **3.12, 3.13 and 3.14** (the `typing` +subscription cache behaves differently per interpreter and this change moves +exactly that behaviour), then `test_eve`, `test_storage`, `test_cartesian`, +`test_examples`; `uv run mypy src/`; `uv run pyright`; `uv run tach check`; +`uv run pre-commit run -a`. One at a time, pytest capped at `-n 4`. + +______________________________________________________________________ + +### PR 2 addendum — the single string moves forward from PR 4 (found during implementation) + +**A gap all four review rounds missed.** Under PR 2 a local dimension's tag becomes its +qualified class name, so `V2EDim = Dimension("V2E", kind=LOCAL)` becomes +`class V2EDim(...)` with tag `mod.V2EDim` — it loses the `"V2E"` spelling that made it +equal to the `FieldOffset` tag and the provider key. **24 of the 38 local-dimension +declarations in `tests/` have a variable name different from their string name**, so PR 2 +by itself breaks constraint A3 (reductions key the provider on the local dimension's name, +`nd_array_field.py:983`, `unroll_reduce.py:47`) and A4 (sparse arguments). The review +checked PR 2 against "conforming programs" and never asked how the test tree would conform +once the tags are qualified. + +**Resolution: PR 4's single-string invariant is pulled forward into PR 2**, with the +existing `V2EDim` class playing the role `V2E.Local` plays later. Every declaration becomes + +```python +class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) +offset_provider = {V2EDim.tag: v2e_table} +``` + +so the FieldOffset tag, the local dimension's tag and the provider key are *one* string +and every path agrees: shifts (PR 1 emits `FieldOffset.value`), reductions and sparse +arguments (the local dim's tag), gtfn's `neighbor_dim`, and DaCe's +`connectivity_identifier(offset_type.tag)`. The #1789 branch becomes dead in PR 2. + +**This reduces total churn rather than adding to it**, provided the references are written +*symbolically*. A use site that says `V2EDim.tag` — not the literal `"tests.….V2EDim"` — +needs no edit in PR 4: PR 4 changes only the declaration, `V2EDim = V2E.Local`, and every +`V2EDim.tag` follows. So the 111 ITIR-level string sites and the provider literals are +rewritten **once, in PR 2, to symbolic `.tag`**, and not again in PR 4. + +The codemod therefore has three jobs, not one: declarations to classes; a `FieldOffset`'s +tag to its *local* dimension's `.tag`; and string provider keys / ITIR offset strings to the +matching symbolic `.tag`. Validate on the central fixtures (`toy_connectivity.py`, +`cases_utils.py`) across every backend before running it tree-wide. + +______________________________________________________________________ + +### PR 3 — `feat[next]: NeighborConnectivity declarations and local dimensions that know their owner` + +**Purely additive**: new concepts next to `FieldOffset`, nothing removed, no +behaviour change, no test churn. + +```python +class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): # §1.4 + owner: ClassVar[type[NeighborConnectivity] | None] = None + max_neighbors: ClassVar[int | None] = None + min_neighbors: ClassVar[int | None] = None + + def __init_subclass__(cls, *, size: int | None = None, **kw): ... + + +class ConnectivityMeta(type): # §1.0 for __hash__ and __getitem__ + @property + def tag(cls) -> Tag: ... + def __call__(cls, *a, **kw) -> NoReturn: ... + + +class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + Local: ClassVar[type[LocalDimensionIndex]] + + def __init_subclass__(cls, *, max_neighbors=None, min_neighbors=None, **kw): ... +``` + +- `__init_subclass__` asserts `"Local" in cls.__dict__` and that it subclasses + `LocalDimensionIndex`, then sets `Local.owner = cls` and copies the counts. A + missing `Local` is a `TypeError` at class creation naming the class. +- **Owner-less** locals: `class LsqUnk(LocalDimensionIndex, size=3)` — `owner is None`, `min == max == size`, never in the provider. ICON4Py's `LsqUnkDim` and + `RBFDimension` need this: they index no table but need sparse storage and + layout. +- Counts are **optional class keywords**, not type parameters (Python has no + integer type parameters and nothing static needs the count). Declared ⇒ a + constraint the table must satisfy. Undeclared ⇒ completed at bind time, from + the table in the JIT flow or from the `NeighborConnectivityType` already passed + through `connectivities=` (`ffront/decorator.py:188-208`) in the AOT flow. + Not static-only because `fvm_nabla_setup.py:99` sizes `V2E` from the atlas + mesh, and ICON skip-value presence is configuration-dependent (`icon.py:130` — + pentagons have skip values on the icosahedron, not the torus). +- **Bind-time validation**, one function replacing constraints A6–A8: shape + `(n, max_neighbors)`, integral dtype, skip values present iff + `min_neighbors < max_neighbors`, `domain[0] is Origin`, `codomain is Codomain`. + +**No provider bridge.** The first draft proposed a dual-keyed +(`Tag | type[NeighborConnectivity]`) provider here. Dropped: the provider is +accessed **directly, not through `get_offset`, at 19 sites in 12 `src/` files** +(`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, +`compiled_program.py`, `pass_manager.py`, …) despite the note at +`common.py:1174`, so a bridge would be both invasive and — since nothing would +exercise class keys — untested. Class-keyed providers land in one place, PR 4. + +**Typing tests**: `typing_probe.py` / `probe_local.py` / `probe_meta_getitem3.py` +become real coverage — `typing_tests/test_next.yaml` cases for +`Field[Dims[V, V2E.Local], float]`, a `TypeVar` bound to `LocalDimensionIndex`, +the negative cross-connectivity case, and the `NC[V, E]`-in-bases case of §1.0; +plus the pyright variants. + +**Acceptance**: full suite green with no behaviour change; new unit tests for +declaration errors, owner wiring, owner-less locals and bind-time validation. + +______________________________________________________________________ + +### PR 4 — `feat[next]: declare connectivities as classes; FieldOffset derived from them` + +**The ordering fix.** Revision 2 put class-keyed providers here and the +declaration migration in PR 6. That cannot be green: `Local.owner` only exists if +the user declared a `NeighborConnectivity` class, and a `FieldOffset` written the +old way (`FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim))`) has no class +to point at — so neither the backend work nor a class-keyed provider has anything +to resolve. The declaration migration must come **first**, and the provider key +must stay a string until the backends are through. + +- The **unstructured** `FieldOffset` is **derived**, not authored: + `FieldOffset.from_connectivity(V2E)` (or `V2E.__gt_offset__()`), which fills + `source = Codomain`, `target = (Origin, V2E.Local)` and — critically — + **`value = V2E.Local.tag`, the *local dimension's* tag, not `V2E.tag`.** +- **Why the local dimension's tag and not the connectivity's.** An earlier draft + used `V2E.tag` and claimed PR 4 was green. It is not: A3 and A4 key the provider + on the **local dimension's** name, at `nd_array_field.py:983` + (`get_offset(provider, axis.value)`, whose in-tree comment is literally + `# assumes offset and local dimension have same name`), `unroll_reduce.py:47` + (`arg.type.offset_type.value`), `gtfn_module.py:95`, `gtir_to_sdfg.py:581, 842`, + `iterator/embedded.py:954, 1519`, and + `gtir_to_sdfg_lambda.py:1371, 1455` (`connectivity_identifier(offset_type.value)`). + Today the `V2EDim = Dimension("V2E")` convention makes that string equal to the + offset tag; PR 4 deletes the convention tree-wide, while the `owner` lookup that + replaces it is PR 5. With `value = V2E.tag` every reduction and every sparse-field + argument would break on embedded, gtfn and DaCe simultaneously — the round-1 + matrix row (`neighbor_sum`, tag≠localdim: embedded/gtfn/DaCe fail) would become + the tree's universal state. + Choosing `V2E.Local.tag` instead makes **all four** of A1, A3, A4 and A5 vacuous + at once, because there is then exactly *one* string and the class produces it. + It also makes the #1789 branch at `itir_to_gtfn_ir.py:181-190` + (`if offset_name != connectivity_type.neighbor_dim.value`) dead already in PR 4. + This is preferable to the alternatives — fusing PR 5 into PR 4, or a transient + `owner`-based fallback inside `get_offset` — because it needs no scaffolding: + the string is simply picked correctly, and PR 5 then removes the dependence on a + string at all. +- **The Cartesian `FieldOffset` constructor stays in PR 4.** An earlier draft said + "every declaration becomes a class", which is wrong: `Ioff`, `Koff` and + `EdgeOffset` (`cases_utils.py:163-169`, e.g. + `Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,))`) and ICON4Py's + `Koff`/`KHalfOff` (`dimension.py:47-48`) are single-target and have no + `NeighborConnectivity` to derive from, and their only remaining consumer — + `as_offset` — does not change until PR 6. Restricting the public constructor to + the single-target form keeps PR 4 green; it disappears with `as_offset` in + PR 6. +- **Providers stay keyed on `Tag`**, now `V2E.Local.tag`. Nothing about the key + *mechanism* changes yet, so the 19 direct-access sites are untouched. A1, A3, A4 + and A5 are all dead at this point — there is one string, and the class produces + it. +- `ts.OffsetType` → `ConnectivityType`, produced by the class + (`type_specifications.py:74` TODO). `type_info.py:637, 858` gate `a(V2E)` + deduction on `ts.OffsetType` and follow. `Field.premap` / `Field.__call__` + unions widen to `Connectivity | type[NeighborConnectivity]` (`common.py:785, 791-794`, `1313-1316`). +- **Test-tree migration lands here**: `toy_connectivity.py`, `cases_utils.py`, + `fvm_nabla_setup.py` are the fixture modules everything imports; 37 + `FieldOffset` sites, 42 `DimensionKind.LOCAL` sites. String provider keys keep + working because they are `cls.tag` — but the tags are now *qualified*, so the + 111 ITIR-level string-offset occurrences in 15 files (`im.shift("V2E")`, + `neighbors("…")`, `OffsetLiteral(value="…")`, string-keyed providers) are + rewritten to `V2E.Local.tag` here rather than in PR 6. +- **ICON4Py**: this is the release-visible declaration change. The migration + script written in PR 2 is extended. + +______________________________________________________________________ + +### PR 5 — `refactor[next]: backends resolve connectivities through the local dimension's owner` + +Where A3, A4 and A5 dissolve and PR 1's skip-matrix entries are removed. Green +while providers are still string-keyed, because a backend goes +`local_dim.owner` → `owner.tag` → the existing lookup: the *identity* question is +answered by the owner pointer, and the key is still a string. + +**The owner-less case must be handled, not assumed away.** At PR 5 `_CONST_DIM` +is still a plain `DimensionIndex(kind=LOCAL)` (it becomes `ConstList` only in +PR 7), and `LsqUnk`-style local axes have `owner is None` by design. Every +converted site reads `getattr(dim, "owner", None)` and falls through when it is +`None`. Two sites already guard by accident — DaCe compares against `_CONST_DIM` +first (`gtir_to_sdfg_lambda.py:1314`) and `unroll_reduce` filters +`offset_type is None` — but `gtfn_module.py:91-98` and `nd_array_field.py:981` +have **no** guard, and ICON4Py never exercises the case +(`test_icon.py:220`), so the gap would not show up downstream. + +`unroll_reduce.py:47` (reads `arg.type.offset_type`, which is the local +`Dimension` — now a `LocalDimensionIndex` carrying `owner`, which is exactly the +back-pointer it lacked), `gtfn_module.py:95, 118, 132`, `itir_to_gtfn_ir.py` +(including the `#1789` `offset_name != neighbor_dim.value` branch at `:181-190`, +which becomes dead and goes), `gtir_to_sdfg.py:581`, +`gtir_to_sdfg_lambda.py:766-770, 1155` (the `Dimension(offset, LOCAL)` synthesis +goes), `nd_array_field.py:981-985` (and its +`# assumes offset and local dimension have same name` comment), +`iterator/embedded.py`, and `runners/roundtrip.py`. + +______________________________________________________________________ + +### PR 6 — `feat[next]!: class-keyed offset providers; remove FieldOffset and the string offset API` + +The only breaking PR, and now the only one that touches the provider key. + +- `OffsetProvider*` become + `Mapping[type[NeighborConnectivity], NeighborTable]`; `get_offset` keys on the + class. **Note the `.owner` hop**: because PR 4 made the IR offset tag the *local + dimension's* tag, `resolve(OffsetLiteral.value)` yields a `LocalDimensionIndex`, + not the connectivity — so the class-key lookup is + `resolve(tag).owner`. (The alternative is to switch the IR tag to `V2E.tag` in + this PR; the `.owner` hop is cheaper and keeps the IR stable.) The **19 direct-access sites in 12 `src/` files** + (`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, + `compiled_program.py`, `pass_manager.py`, …) are converted here, and + `common.py:1174`'s "all accesses should go through `get_offset`" either becomes + true or the note goes. `hash_offset_provider_items_by_id` and the + `fingerprinting` dict handling already tolerate class keys once §1.0's + `__hash__` is in place. +- **Removals**: `FieldOffset` entirely (both forms); `runtime.Offset` as its base + (`fbuiltins.py:467-470` TODO); `iterator/runtime.offset("...")` (12 sites in 6 + files, plus `tracing.py:161-162`); the `V2EDim`-next-to-`V2E` convention; + `embedded/context.py` string plumbing; the `gt4py.next.__init__` exports at + `:47, 140`. +- **`as_offset` changes in the same PR.** It is why the Cartesian `FieldOffset` + form cannot go alone: `ffront/experimental.py:17` + + `type_deduction.py:956-967` require one. New signature + `as_offset(KDim, field)`. Used in 5 test modules, the `Ioff`/`Koff`/`EdgeOffset` + fixtures at `cases_utils.py:163-169`, and **40 non-test call sites in + ICON4Py**. +- `transform_utils.py:50-77` and `past_to_itir.py:77` deduce grid type from the + provider and follow. +- **Accepted double churn**: the ~26 `offset_provider={...}` literals are rewritten + twice — `{V2E.Local.tag: t}` in PR 4, `{V2E: t}` here. The alternative is fusing PR 4 + and PR 6, which loses the green boundary. The ITIR string sites do *not* churn + twice: `im.shift(V2E.Local.tag)` written in PR 4 stays correct. + +**ADR 0029**: the connectivities-as-types record — `FieldOffset` removed, the +class-keyed provider, superseding the `FieldOffset` part of ADR 0019. + +**Breaking-change communication**: the PR title carries the Conventional Commits +`!` marker and the ADR records the removal; the changelog entry is written by the +release PR, not here (see PR 1). No deprecation window — an explicit decision: +ICON4Py's provider keys are bare names, so no import-based shim could have +resolved them. + +**Acceptance**: full suite; `test_fvm_nabla` and `test_icon_like_scan` are the +integration canaries. A before/after run of one gtfn and one DaCe program +checking generated-code equivalence modulo names. + +### PR 7 — `refactor[next]: ConstList replaces _CONST_DIM; AxisLiteral drops kind` + +- `_CONST_DIM` (`iterator/embedded.py:220` and + `dace/lowering/gtir_to_sdfg_lambda.py` — two separate declarations, each + internally consistent; 14 references in total) becomes the owner-less + `ConstList(LocalDimensionIndex, size=1)`, generalizing the magic name from size + 1 to size *n*. +- `AxisLiteral.kind` removed (`iterator/ir.py:93` TODO), now that every tag + resolves to a class carrying its kind. + +______________________________________________________________________ + +### PR 8 — `refactor[next]: MultiDimensionIndex and typed embedded positions` + +- `MultiDimensionIndex[D: DimensionIndex, *Ls]` as the index type of a sparse + position and the domain index of a `NeighborTable`. `*Ls` is unconstrained + because `TypeVarTuple` cannot carry a bound; `__init_subclass__` checks at + runtime what the checker cannot. +- `iterator/embedded.py` positions keyed by dimension types instead of name + strings (`embedded.py:574-576`, `597-616`, `941-950`); `SparseTag` removed. + Constraint A9 dissolves. +- Nothing else depends on this; it is last for that reason. + +______________________________________________________________________ + +## 3. Constraint ledger + +| # | Constraint | Retired by | +| -------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | ------------------------------------------------------- | +| A1 | `FieldOffset.value` == provider key | PR 4 (one declaration produces both) | +| — | *all four string-equality constraints below become vacuous in PR 4*, because the class emits a single string (`V2E.Local.tag`); PR 5 removes the dependence on a string at all | PR 4 / PR 5 | +| A2 | Python variable name == provider key | **PR 1** | +| A3 | local dim name == provider key (reductions) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | +| A4 | local dim name == provider key (sparse args) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | +| A5 | `FieldOffset.value` == local dim name | PR 4 (one declaration) | +| A6 | `target[-1]` == connectivity `neighbor_dim` | PR 3 (bind-time check) | +| A7 | `FieldOffset.source` == `codomain` | PR 3 (bind-time check) | +| A8 | `target[0]` == `domain[0]` | PR 3 (bind-time check) | +| A9 | dim name is the iterator-position dict key | PR 8 | +| A10 | dim name round-trips through `AxisLiteral` | structural; PR 2 qualifies it, PR 7 drops `kind` | +| F5/F8/F9 | codegen name formats | PR 2 (`codegen_name` + inverse) | +| S6 | `as_offset` needs a Cartesian `FieldOffset` | PR 6 | +| — | string-keyed provider | PR 6 | + +## 4. Risks + +1. **PR 2 size, and it is not purely mechanical.** ~150 files of migration *plus* + a name-mangling layer and `Staggered[D]`. It cannot be reviewed as "the + #2844 diff plus identity". Mitigation: land the mangling layer and + `Staggered[D]` as reviewable commits *within* the PR, ordered before the + sweep, so the mechanical part is a separate commit. +2. **Function-local dummy dimensions.** 133 declarations in 15 test files must + move to module level. Under nominal identity, two same-named locals that were + silently the same dimension become distinct — each resulting failure is a + real finding, not churn. +3. **`typing` subscription caching.** Under `(name, kind)` equality, + `Field[Dims[I]] is Field[Dims[I2]]` aliases for two distinct same-named + classes — a known residual of the #2844 design, and an argument *for* this + stack. The cache behaves differently per interpreter: verify on 3.12, 3.13 + and 3.14 **via nox**, not `uv run pytest`. +4. **Fingerprint/cache invalidation.** Moving a declaration between modules now + invalidates compiled artifacts. Intended; CHANGELOG + ADR line. +5. **Naming not yet converged** with `havogt/dependent-local-dimensions` + (`Origin`/`Codomain` vs `source_dim`/`neighbor_dim`; `Local` vs `Dim`; + `min_neighbors` vs `has_skip_values`). PR 3 fixes public names. Converge + before PR 3 is *opened*. +6. **`V2E.Local` vs `Local[V2E]`.** The chain proposals' encodings subscript + `Local`, and a `TypeVar` cannot be subscripted for a nested attribute. Their + semantics are unaffected; their static encoding needs rewriting to `C.Local` + plus a protocol for the generic hop-stack case. Knowledge-repo concern, not a + gt4py blocker. +7. **CSCS GPU CI is flaky and opaque.** All jobs failing at the same second means + infrastructure; `cscs-ci run default` as a PR comment reruns it. #2844's CI is + green except that job. +8. **Two mangling passes, two PRs.** `codegen_name` is applied to dimension names + in PR 2 and to offset names in PR 4, at ~19 sites total, several of which + (`Sym`/`SymRef` construction, DaCe array names) fail *loudly* and several of + which (C++ emission, `name.lower()`) fail only in the generated artifact. + Both PRs need a test that a qualified tag survives a real gtfn and a real + DaCe compile, not just lowering. +9. **`resolve()` on the inference hot path** (`inference.py:464`, once per + `AxisLiteral`). Must be memoized from the start, and the memo must be keyed + so a reloaded module does not return a stale class. + +## 4b. Work the earlier drafts did not mention + +- **Public exports.** `gt4py.next.__init__` must export `NeighborConnectivity`, + `LocalDimensionIndex`, `Staggered` and `resolve` (PR 2 for the dimension half, + PR 3 for the connectivity half), and drop `FieldOffset` / `offset` at `:47, 140` + in PR 6. +- **`type_translation.from_value(V2E)`** works only because the + `hasattr(value, "__gt_type__")` branch at `type_translation.py:328` is tested + *before* the `DimensionMeta` branch. That ordering is load-bearing under this + design and currently untested — PR 3 adds a unit test pinning it. +- **`pyright` is not yet a dependency.** `uv run pyright` appears throughout §5 + but pyright is absent from `pyproject.toml` on `main`; #2845 is what adds it. + PR 2 must explicitly fold in #2845's `typing_exports` / pyright dependency-group + change, or §5's pyright step is not runnable. +- **`test_examples` belongs to PR 4 too.** §5 lists it for PR 2 and PR 6; the + docs and notebooks use `FieldOffset`, so PR 4's declaration migration touches + them and must run it. + +## 5. Verification + +Per PR, in this order, **one at a time** on the shared machine, pytest capped at +`-n 4`: + +``` +uv run pre-commit run -a # ruff, mypy, tach, license headers +uv run pyright # static checks the mypy plugin no longer fakes +uv run nox -s "test_next-3.12(...)" # then 3.13, 3.14 for PR 2 +uv run nox -s test_eve test_storage test_cartesian test_examples # PR 2, PR 6 +``` + +Test-first where behaviour changes, per AGENTS.md: PR 1's regression matrix, PR +3's declaration-error and bind-validation units, and PR 4's provider-key tests +are written before the implementation they cover. + +## 6. Feedback owed to the knowledge-repo note + +- `NeighborConnectivity` is not a `Connectivity`; the sketch's base line is wrong + (§1.1) — this closes Open Q6. +- `DimensionBaseIndex` should be dropped; `LocalDimensionIndex` subclasses + `DimensionIndex` (§1.4). +- Open Q2 is closed: metaclass discovery, base `ClassVar` *annotation*, explicit + nested class — with the precise limit of what that buys statically (§1.2). +- `Staggered[D]` is not a late step; it is a precondition for removing the + name-keyed registry (§1.5) — and it cannot be a PEP 695 generic, needs an + identity-keyed intern cache, and needs a narrow `copyreg`. The note's claim + that `copyreg` disappears entirely is therefore too strong. +- The note's "five name spaces" analysis should record that under qualified tags + there are **two** dotted name spaces reaching codegen — dimension tags and + offset tags — each needing its own mangling pass (§1.3(b) and (c)). +- The note's §Staging step 2 ("`NeighborConnectivity` … object-keyed provider; + `FieldOffset` and string keys removed outright") bundles three changes that + must be separated to stay green: the declaration migration has to precede the + backend work (because `Local.owner` only exists once classes are declared), and + the provider *key* has to stay a string until the backends resolve through the + owner. See PR 4/5/6. +- The staging in the note's §Staging (steps 0–8) is superseded by §2 here; in + particular step 4 ("backends, one file at a time") cannot follow step 2, since + the provider key change and the backend lookups are separable but the mangling + layer is needed at the *dimension* step. + +## 7. Open, non-blocking + +- Naming convergence (risk 5). +- Whether a same-`__name__` collision warning is useful or noise (note Q5). +- Whether interactive `__main__` should be detected with a fallback to in-process + compilation and a warning, or merely documented (note Q3). Plan assumes + documented. +- `AxisLiteral.dim` instead of `AxisLiteral.value` — a follow-up after PR 8. + +______________________________________________________________________ + +## 8. As implemented (PRs 3–6), and where it deviates + +Recorded while implementing; the sections above are the plan as reviewed. + +- **PR 3 was made usable in the DSL**, not just declarative: `ConnectivityMeta.__gt_type__` + gives the offset type, `V2E[i]` and `a(V2E)` work in embedded and on all backends, and + `V2E.Local` in DSL code types as the local dimension (one `visit_Attribute` case). No backend + change was needed, because PR 2 had already made the offset tag, the local dimension's tag and + the provider key one string. Also in PR 3: fingerprinting of a declaration by its dimensions and + counts, `FieldOffset.Local`, and `check_neighbor_table` (explicit until PR 6). +- **Shared local dimensions (not in the plan).** The review of PR 4 found that ICON4Py's flattened + sparse offsets (`C2CE`, `E2ECV`, `E2EC`, `C2CEC`) use another connectivity's local dimension, so + the one-owner rule would have made them unmigratable once `FieldOffset` is gone. A declaration + can adopt an owned local dimension (`Local = C2E.Local`); the owner stays `C2E`, and the sharer + is named in the IR by its own tag: `ConnectivityMeta.offset_tag` is `Local.tag` for the owner + and `cls.tag` for a sharer. This reintroduces an offset tag that differs from its local + dimension's tag — exactly PR 1's `tag != local dim` cells. +- **PR 5 was re-scoped** from "backends resolve through `owner`" to "backends support + connectivities sharing a local dimension". With the single string of PR 2 and class-keyed + providers normalized to tags (PR 6), resolving through `owner` would change nothing; what the + backends did lack was a way to find the table over a local dimension that is not keyed by its + tag. `common.connectivity_key_over(provider, local_dim)` does that (the local dimension's tag if + bound, else any table whose neighbor dimension it is), and replaces the local-tag lookups in + embedded reductions, `unroll_reduce`, gtfn sparse arguments, iterator-embedded sparse lists and + positions, and DaCe (`make_field`, reductions, `map_list`, local-dimension array sizes). All of + PR 1's `uses_offset_tag_differing_from_local_dim*` markers are gone in PR 5. +- **PR 6 keeps the internal provider tag-keyed.** Users key providers by the class; every entry + point (`Program.__call__`, `FieldOperator.__call__`, `compile`, `CompilationOptions`, + `embedded.context.update`, iterator `fendef`) normalizes with `as_tag_keyed_offset_provider` + (class → `offset_tag`; a bare, undotted string key is rejected as the removed `FieldOffset` + spelling). Tables are checked against their declarations once per compiled variant + (`check_offset_provider` in `CompiledProgramsPool._compile_variant`) and on embedded calls, not + on every call. The 19 direct-access sites of §PR 6 therefore needed no change, and the IR keeps + naming connectivities by string, consistently with §1.3. +- **`as_offset(KDim, field)`** takes a dimension in PR 6; Cartesian `FieldOffset`s had no other + use (Cartesian shifts are `Dim + i` since before this stack). +- **Not done**: `ts.OffsetType` is not renamed to `ConnectivityType` (it now only types + connectivity declarations and `as_offset`); `iterator.runtime.offset("...")` stays as the + iterator-level API, since the IR names offsets by string. diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 08c37acb03..8def73342a 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -231,7 +231,7 @@ class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): class Local(gtx.LocalDimensionIndex): ... ``` -Note that the declaration does not contain the actual connectivity table, that's provided through an _offset provider_, keyed by the local dimension's `tag`: +Note that the declaration does not contain the actual connectivity table, that's provided through an _offset provider_, a dictionary from connectivity declarations to tables: ```{code-cell} ipython3 E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2C.Local], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) @@ -250,7 +250,7 @@ def nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64]) -> gtx. def run_nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): nearest_cell_to_edge(cell_values, out=out) -run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider}) +run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2C: E2C_offset_provider}) print("0th adjacent cell's value: {}".format(edge_values.asnumpy())) ``` @@ -277,7 +277,7 @@ def sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64]) -> gtx.Field[D def run_sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): sum_adjacent_cells(cells, out=out) -run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider}) +run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2C: E2C_offset_provider}) print("sum of adjacent cells: {}".format(edge_values.asnumpy())) ``` @@ -441,7 +441,7 @@ result_pseudo_lap = gtx.as_field([CellDim], np.zeros(shape=(6,))) run_pseudo_laplacian(cell_values, edge_weight_field, result_pseudo_lap, - offset_provider={E2C.Local.tag: E2C_offset_provider, C2E.Local.tag: C2E_offset_provider}) + offset_provider={E2C: E2C_offset_provider, C2E: C2E_offset_provider}) print("pseudo-laplacian: {}".format(result_pseudo_lap.asnumpy())) ``` diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb index 21bf2d25d8..0d98021f37 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb @@ -126,7 +126,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.Local.tag: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb index 86c8d33ac7..5176f48285 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb @@ -131,7 +131,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.Local.tag: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb index fb2282ab22..77d1ecf296 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb @@ -123,7 +123,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.Local.tag: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb index 43196507ce..b90958ea63 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb @@ -136,7 +136,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.Local.tag: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb index b99c6f6d3f..ed332054f8 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb @@ -147,7 +147,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.Local.tag: v2e_connectivity},\n", + " offset_provider={V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb index de040ccb93..0f76b5a6b1 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb @@ -152,7 +152,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.Local.tag: v2e_connectivity},\n", + " offset_provider={V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb index 174699e350..b557d404ea 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb @@ -293,10 +293,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.Local.tag: c2e_connectivity,\n", - " V2E.Local.tag: v2e_connectivity,\n", - " E2V.Local.tag: e2v_connectivity,\n", - " E2C.Local.tag: e2c_connectivity,\n", + " C2E: c2e_connectivity,\n", + " V2E: v2e_connectivity,\n", + " E2V: e2v_connectivity,\n", + " E2C: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb index 81836edbd5..7e1a7ecc4c 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb @@ -314,10 +314,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.Local.tag: c2e_connectivity,\n", - " V2E.Local.tag: v2e_connectivity,\n", - " E2V.Local.tag: e2v_connectivity,\n", - " E2C.Local.tag: e2c_connectivity,\n", + " C2E: c2e_connectivity,\n", + " V2E: v2e_connectivity,\n", + " E2V: e2v_connectivity,\n", + " E2C: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb index edd65ac9e5..26973ad5b4 100644 --- a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb @@ -169,7 +169,7 @@ " kappa,\n", " dt,\n", " out=(divergence_gt4py_1, divergence_gt4py_2),\n", - " offset_provider={E2C2V.Local.tag: e2c2v_connectivity, V2E.Local.tag: v2e_connectivity},\n", + " offset_provider={E2C2V: e2c2v_connectivity, V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py_1.asnumpy(), divergence_ref_1)\n", diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index f16f1560bc..2cde0190ff 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -321,7 +321,7 @@ " nearest_cell_to_edge(cell_field, out=edge_field)\n", "\n", "\n", - "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", + "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={E2C: E2C_offset_provider})\n", "\n", "print(\"0th adjacent cell's value: {}\".format(edge_field.asnumpy()))" ] @@ -396,7 +396,7 @@ " sum_adjacent_cells(cell_field, out=edge_field)\n", "\n", "\n", - "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", + "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={E2C: E2C_offset_provider})\n", "\n", "print(\"sum of adjacent cells: {}\".format(edge_field.asnumpy()))" ] diff --git a/scripts/python/migrate_connectivities.py b/scripts/python/migrate_connectivities.py new file mode 100644 index 0000000000..d8862873b5 --- /dev/null +++ b/scripts/python/migrate_connectivities.py @@ -0,0 +1,336 @@ +#!/usr/bin/env -S uv run -q --frozen --isolated --python 3.12 --group scripts python3 +# +# 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 + +""" +Migrate gt4py.next user code to dimension and connectivity classes (ADRs 0028, 0029). + +Rewrites module-level declarations and the uses of Cartesian offsets: + + KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) + E2CDim = gtx.Dimension("E2C", gtx.DimensionKind.LOCAL) + E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim, E2CDim)) + Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) + ... a(Koff[1]) ... as_offset(Koff, k_field) ... + +becomes + + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class E2CDim(gtx.LocalDimensionIndex): ... + class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + Local = E2CDim + ... a(KDim + 1) ... as_offset(KDim, k_field) ... + +A connectivity adopts its existing local dimension (`Local = E2CDim`), so the names already used +for local dimensions keep working, and so do offsets that share a local dimension (`C2CE` +with `C2EDim`). What cannot be rewritten from the source alone is reported instead: offset +providers keyed by strings, which become keyed by the connectivity (`{E2C: table}`), `.value` +on a dimension (now `.tag`), and `isinstance` checks against `Dimension`. + +The output is not formatted; run `ruff format` on the changed files afterwards. +""" + +from __future__ import annotations + +import ast +import dataclasses +import difflib +import pathlib +import re +from collections.abc import Iterable, Iterator + +import typer + + +cli = typer.Typer(no_args_is_help=True, name="migrate-connectivities", help=__doc__) + + +@dataclasses.dataclass(frozen=True) +class Edit: + """Replace source lines `[start, end)` (0-based) with `text`.""" + + start: int + end: int + text: str + + +@dataclasses.dataclass +class Module: + path: pathlib.Path + source: str + tree: ast.Module + edits: list[Edit] = dataclasses.field(default_factory=list) + notes: list[str] = dataclasses.field(default_factory=list) + + @property + def lines(self) -> list[str]: + return self.source.splitlines(keepends=True) + + def segment(self, node: ast.AST) -> str: + segment = ast.get_source_segment(self.source, node) + assert segment is not None + return segment + + def note(self, node: ast.AST, message: str) -> None: + self.notes.append(f"{self.path}:{node.lineno}: {message}") + + +def _callee_name(call: ast.Call) -> tuple[str, str] | None: + """`(prefix, name)` of a call to `[prefix.]name`, e.g. `("gtx.", "Dimension")`.""" + match call.func: + case ast.Name(id=name): + return "", name + case ast.Attribute(value=value, attr=name): + return f"{ast.unparse(value)}.", name + return None + + +def _keyword(call: ast.Call, name: str, position: int) -> ast.expr | None: + for keyword in call.keywords: + if keyword.arg == name: + return keyword.value + return call.args[position] if len(call.args) > position else None + + +def _declarations(module: Module) -> Iterator[tuple[ast.Assign, str, ast.Call, str, str]]: + """Module-level `name = [prefix.]Dimension(...)` / `FieldOffset(...)` statements.""" + for statement in module.tree.body: + if ( + isinstance(statement, ast.Assign) + and len(statement.targets) == 1 + and isinstance(statement.targets[0], ast.Name) + and isinstance(statement.value, ast.Call) + and (callee := _callee_name(statement.value)) is not None + and callee[1] in ("Dimension", "FieldOffset") + ): + yield statement, statement.targets[0].id, statement.value, *callee + + +def _replace(module: Module, statement: ast.stmt, text: str) -> None: + assert statement.end_lineno is not None + module.edits.append(Edit(statement.lineno - 1, statement.end_lineno, text)) + + +def _migrate_declarations(module: Module, cartesian: dict[str, str]) -> None: + for statement, name, call, prefix, kind_of_call in _declarations(module): + if kind_of_call == "Dimension": + kind = _keyword(call, "kind", 1) + kind_src = module.segment(kind) if kind is not None else None + if kind_src is not None and kind_src.endswith(".LOCAL"): + text = f"class {name}({prefix}LocalDimensionIndex): ...\n" + elif kind_src is None: + text = f"class {name}({prefix}DimensionIndex): ...\n" + else: + text = f"class {name}({prefix}DimensionIndex, kind={kind_src}): ...\n" + _replace(module, statement, text) + continue + + source, target = _keyword(call, "source", 1), _keyword(call, "target", 2) + if source is None or not isinstance(target, ast.Tuple): + module.note(statement, f"'{name}': unrecognized 'FieldOffset' arguments, not migrated.") + continue + if len(target.elts) == 2: + origin, local = (module.segment(element) for element in target.elts) + text = ( + f"class {name}({prefix}NeighborConnectivity[{origin}, {module.segment(source)}]):\n" + f" Local = {local}\n" + ) + _replace(module, statement, text) + elif len(target.elts) == 1 and module.segment(target.elts[0]) == module.segment(source): + # A Cartesian offset has no declaration any more: `Off[i]` is `Dim + i`. + cartesian[name] = module.segment(source) + _replace(module, statement, "") + else: + module.note(statement, f"'{name}': a cross-dimension offset has no class equivalent.") + + +class _CartesianUses(ast.NodeVisitor): + """Rewrite the uses of removed Cartesian offsets, and note the ones it cannot.""" + + def __init__(self, module: Module, cartesian: dict[str, str]) -> None: + self.module = module + self.cartesian = cartesian + #: (line, start column, end column, replacement) of single-line expression rewrites + self.rewrites: list[tuple[int, int, int, str]] = [] + self.handled: set[int] = set() + + def _rewrite(self, node: ast.expr, text: str) -> None: + assert node.end_lineno is not None and node.end_col_offset is not None + if node.lineno != node.end_lineno: + self.module.note(node, f"multi-line expression; rewrite by hand as '{text}'.") + return + self.rewrites.append((node.lineno - 1, node.col_offset, node.end_col_offset, text)) + + def _dimension_of(self, node: ast.expr) -> str | None: + """The dimension replacing `Off` or `module.Off`, qualified like the offset was.""" + match node: + case ast.Name(id=name) if name in self.cartesian: + return self.cartesian[name] + case ast.Attribute(value=value, attr=name) if name in self.cartesian: + return f"{self.module.segment(value)}.{self.cartesian[name]}" + return None + + def visit_Subscript(self, node: ast.Subscript) -> None: + if (dim := self._dimension_of(node.value)) is not None: + match node.slice: + case ast.Constant(value=int() as index): + text = f"{dim} + {index}" if index >= 0 else f"{dim} - {-index}" + case ast.UnaryOp(op=ast.USub(), operand=ast.Constant(value=int() as index)): + text = f"{dim} - {index}" + case _: + text = f"{dim} + ({self.module.segment(node.slice)})" + self._rewrite(node, text) + self.handled.add(id(node.value)) + self.generic_visit(node) + + def visit_Name(self, node: ast.Name) -> None: + if id(node) not in self.handled and (dim := self._dimension_of(node)) is not None: + # e.g. `as_offset(Koff, field)`, which now takes the dimension + self._rewrite(node, dim) + + def visit_Attribute(self, node: ast.Attribute) -> None: + if id(node) not in self.handled and (dim := self._dimension_of(node)) is not None: + self._rewrite(node, dim) + else: + self.generic_visit(node) + + +def _migrate_cartesian_uses(module: Module, cartesian: dict[str, str]) -> None: + if not cartesian: + return + visitor = _CartesianUses(module, cartesian) + import_lines: set[int] = set() + for statement in ast.walk(module.tree): + if isinstance(statement, ast.ImportFrom) and any( + alias.name in cartesian for alias in statement.names + ): + import_lines.update(range(statement.lineno - 1, statement.end_lineno or 0)) + names = [ + cartesian.get(alias.name, alias.name) if alias.asname is None else alias.name + for alias in statement.names + ] + unique = list(dict.fromkeys(names)) + _replace( + module, + statement, + " " * statement.col_offset + + f"from {'.' * statement.level}{statement.module or ''} import {', '.join(unique)}\n", + ) + for statement in module.tree.body: + if not isinstance(statement, (ast.Import, ast.ImportFrom)): + visitor.visit(statement) + + lines = module.lines + for line, start, end, text in sorted(visitor.rewrites, reverse=True): + if line in import_lines: + continue + lines[line] = lines[line][:start] + text + lines[line][end:] + module.source = "".join(lines) + + +_PROVIDER_KEY_RE = re.compile(r"""(?P["'])(?P[A-Za-z_]\w*)(?P=quote)\s*:""") + + +def _report(module: Module, offset_names: set[str], dimension_names: set[str]) -> None: + for number, line in enumerate(module.source.splitlines(), start=1): + for match in _PROVIDER_KEY_RE.finditer(line): + if match["name"] in offset_names: + module.notes.append( + f"{module.path}:{number}: offset-provider key '{match['name']}' is keyed by" + f" the connectivity class now, e.g. '{{{match['name']}: table}}'." + ) + for node in ast.walk(module.tree): + match node: + case ast.Attribute(value=ast.Name(id=name), attr="value") if name in dimension_names: + module.note(node, f"'{name}.value': a dimension's name is '{name}.tag' now.") + case ast.Call( + func=ast.Name(id="isinstance"), args=[_, ast.Attribute(attr="Dimension")] + ): + module.note( + node, + "'isinstance(..., Dimension)': a dimension is a class now; use" + " 'isinstance(obj, gt4py.next.common.DimensionMeta)'.", + ) + + +def _apply(module: Module) -> str: + lines = module.lines + for edit in sorted(module.edits, key=lambda edit: edit.start, reverse=True): + lines[edit.start : edit.end] = [edit.text] if edit.text else [] + return "".join(lines) + + +def migrate(sources: dict[pathlib.Path, str]) -> tuple[dict[pathlib.Path, str], list[str]]: + """Migrate the given modules; return their new sources and what is left to do by hand.""" + modules = [ + Module(path=path, source=source, tree=ast.parse(source)) for path, source in sources.items() + ] + cartesian: dict[str, str] = {} + offset_names: set[str] = set() + dimension_names: set[str] = set() + for module in modules: + for _, name, _, _, kind_of_call in _declarations(module): + (dimension_names if kind_of_call == "Dimension" else offset_names).add(name) + _migrate_declarations(module, cartesian) + + results: dict[pathlib.Path, str] = {} + notes: list[str] = [] + for module in modules: + _report(module, offset_names, dimension_names) + migrated = _apply(module) + # Cartesian uses are rewritten on the migrated text, re-parsed, so that line numbers + # refer to what the declaration edits left. + second = Module(path=module.path, source=migrated, tree=ast.parse(migrated)) + _migrate_cartesian_uses(second, cartesian) + results[module.path] = _apply(second) + notes += module.notes + second.notes + return results, notes + + +def _python_files(paths: Iterable[pathlib.Path]) -> Iterator[pathlib.Path]: + for path in paths: + if path.is_dir(): + yield from sorted(path.rglob("*.py")) + else: + yield path + + +@cli.command() +def run( + paths: list[pathlib.Path], + write: bool = typer.Option(False, "--write", help="Rewrite files instead of printing a diff."), +) -> None: + """Migrate the Python files in PATHS (directories are searched recursively).""" + sources = {path: path.read_text() for path in _python_files(paths)} + results, notes = migrate(sources) + for path, new in results.items(): + if new == sources[path]: + continue + if write: + path.write_text(new) + typer.echo(f"migrated {path}") + else: + typer.echo( + "".join( + difflib.unified_diff( + sources[path].splitlines(keepends=True), + new.splitlines(keepends=True), + fromfile=str(path), + tofile=str(path), + ) + ) + ) + if notes: + typer.echo("\nLeft to migrate by hand:", err=True) + for note in notes: + typer.echo(f" {note}", err=True) + + +if __name__ == "__main__": + cli() diff --git a/scripts/tests/python/test_migrate_connectivities.py b/scripts/tests/python/test_migrate_connectivities.py new file mode 100644 index 0000000000..a8d48ad26d --- /dev/null +++ b/scripts/tests/python/test_migrate_connectivities.py @@ -0,0 +1,116 @@ +# +# 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 +# + +"""Tests for the ``migrate_connectivities`` dev script.""" + +from __future__ import annotations + +import pathlib +import textwrap + +import migrate_connectivities + + +DIMENSIONS = textwrap.dedent( + """\ + import gt4py.next as gtx + + KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) + EdgeDim = gtx.Dimension("Edge") + CellDim = gtx.Dimension("Cell") + CEDim = gtx.Dimension("CE") + E2CDim = gtx.Dimension("E2C", gtx.DimensionKind.LOCAL) + C2EDim = gtx.Dimension("C2E", gtx.DimensionKind.LOCAL) + E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim, E2CDim)) + C2E = gtx.FieldOffset("C2E", source=EdgeDim, target=(CellDim, C2EDim)) + C2CE = gtx.FieldOffset("C2CE", source=CEDim, target=(CellDim, C2EDim)) + Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) + """ +) + +STENCIL = textwrap.dedent( + """\ + import gt4py.next as gtx + from gt4py.next.ffront.experimental import as_offset + + from pkg import dimension as dims + from pkg.dimension import E2C, KDim, Koff + + + def stencil(a, k): + b = a(Koff[1]) + a(Koff[-1]) + a(Koff[k]) + a(dims.Koff[1]) + return b(as_offset(Koff, k)) + b(as_offset(dims.Koff, k)) + + + def run(prog, grid): + prog(offset_provider={"E2C": grid.e2c, "Koff": KDim}) + return KDim.value + """ +) + + +def _migrate(**sources: str) -> tuple[dict[str, str], list[str]]: + results, notes = migrate_connectivities.migrate( + {pathlib.Path(name): source for name, source in sources.items()} + ) + return {str(path): source for path, source in results.items()}, notes + + +def test_declarations(): + results, _ = _migrate(dimension=DIMENSIONS) + + assert results["dimension"] == textwrap.dedent( + """\ + import gt4py.next as gtx + + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class EdgeDim(gtx.DimensionIndex): ... + class CellDim(gtx.DimensionIndex): ... + class CEDim(gtx.DimensionIndex): ... + class E2CDim(gtx.LocalDimensionIndex): ... + class C2EDim(gtx.LocalDimensionIndex): ... + class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + Local = E2CDim + class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + Local = C2EDim + class C2CE(gtx.NeighborConnectivity[CellDim, CEDim]): + Local = C2EDim + """ + ) + + +def test_declarations_run(): + results, _ = _migrate(dimension=DIMENSIONS) + namespace: dict = {"__name__": "migrated_dimension"} + exec(results["dimension"], namespace) + + assert namespace["C2E"].Local is namespace["C2EDim"] + assert namespace["C2EDim"].owner is namespace["C2E"] + # `C2CE` shares `C2E`'s local dimension and is named by its own tag + assert namespace["C2CE"].Local is namespace["C2EDim"] + assert namespace["C2CE"].offset_tag == namespace["C2CE"].tag + + +def test_cartesian_offset_uses_across_modules(): + results, _ = _migrate(dimension=DIMENSIONS, stencil=STENCIL) + + stencil = results["stencil"] + assert "from pkg.dimension import E2C, KDim\n" in stencil + assert "a(KDim + 1) + a(KDim - 1) + a(KDim + (k)) + a(dims.KDim + 1)" in stencil + assert "b(as_offset(KDim, k)) + b(as_offset(dims.KDim, k))" in stencil + assert "Koff" not in stencil.replace('"Koff"', "") + + +def test_what_is_left_is_reported(): + _, notes = _migrate(dimension=DIMENSIONS, stencil=STENCIL) + + assert any("offset-provider key 'E2C'" in note for note in notes) + assert any("offset-provider key 'Koff'" in note for note in notes) + assert any("'KDim.value'" in note for note in notes) diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index 94690b7d0e..1fd3b4d806 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -50,7 +50,6 @@ from .ffront import fbuiltins from .ffront.decorator import field_operator, program, scan_operator from .ffront.fbuiltins import ( - FieldOffset, IndexType, abs, # noqa: A004 # shadowing arccos, @@ -149,7 +148,6 @@ "as_field", "as_connectivity", # from ffront - "FieldOffset", "field_operator", "program", "scan_operator", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 5ffc813f8c..3128b6e535 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -390,6 +390,15 @@ def resolve(tag: Tag) -> Dimension: ) return owner[resolve(match["base"])] # type: ignore[index] # a StaggeredMeta, checked + obj = _import_qualified_name(tag) + if not isinstance(obj, DimensionMeta): + raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") + return cast(Dimension, obj) + + +@functools.cache +def _import_qualified_name(tag: Tag) -> Any: + """Import the object a dotted qualified name refers to; see `resolve`.""" parts = tag.split(".") for split in range(len(parts), 0, -1): try: @@ -404,9 +413,7 @@ def resolve(tag: Tag) -> Dimension: f"Cannot resolve dimension tag '{tag}': '{'.'.join(parts[:split])}' has" f" no attribute '{attr}'." ) from ex - if not isinstance(obj, DimensionMeta): - raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") - return cast(Dimension, obj) + return obj raise ValueError( f"Cannot resolve dimension tag '{tag}': no importable module prefix. A dimension" " referenced from the IR must be declared at module level in an importable module." @@ -1036,9 +1043,7 @@ def asnumpy(self) -> np.ndarray: ... def as_scalar(self) -> core_defs.ScalarT: ... @abc.abstractmethod - def premap( - self, index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity] - ) -> Field: ... + def premap(self, index_field: Connectivity | type[NeighborConnectivity]) -> Field: ... @abc.abstractmethod def restrict(self, item: AnyIndexSpec) -> Self: ... @@ -1046,8 +1051,8 @@ def restrict(self, item: AnyIndexSpec) -> Self: ... @abc.abstractmethod def __call__( self, - index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], - *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + index_field: Connectivity | type[NeighborConnectivity], + *args: Connectivity | type[NeighborConnectivity], ) -> Field: ... @abc.abstractmethod @@ -1432,8 +1437,16 @@ def is_neighbor_table(obj: Any) -> TypeGuard[NeighborTable]: OffsetProviderTypeElem: TypeAlias = NeighborConnectivityType # Note: `OffsetProvider` and `OffsetProviderType` should not be accessed directly, # use the `get_offset` and `get_offset_type` functions instead. +#: Neighbor tables keyed by the connectivity's `offset_tag`, which is how the IR names it. OffsetProvider: TypeAlias = Mapping[Tag, OffsetProviderElem] OffsetProviderType: TypeAlias = Mapping[Tag, OffsetProviderTypeElem] +#: An offset provider as users write it: keyed by `NeighborConnectivity` declarations (or, at the +#: IR level, by tags). The entry points of a program normalize it to an `OffsetProvider` with +#: `as_tag_keyed_offset_provider`, so everything below them sees tags only. +#: NOTE: `Any` keys, since `Mapping` is invariant in its key type: a tag-keyed provider would not +#: be a `Mapping[type[NeighborConnectivity] | Tag, ...]`. Keys are checked at runtime instead. +OffsetProviderLike: TypeAlias = Mapping[Any, OffsetProviderElem] +OffsetProviderTypeLike: TypeAlias = Mapping[Any, OffsetProviderTypeElem] def is_offset_provider(obj: Any) -> TypeGuard[OffsetProvider]: @@ -1464,8 +1477,12 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid """ # TODO(havogt): Once we have a custom class for `OffsetProvider`, we can absorb this functionality into it. if offset_tag not in offset_provider: - raise KeyError(f"Offset '{offset_tag}' not found in offset provider.") - return offset_provider[offset_tag] # TODO return a valid dimension + raise KeyError( + f"Connectivity '{offset_tag}' not found in the offset provider, which has" + f" {sorted(map(str, offset_provider))}. Offset providers are keyed by" + " 'NeighborConnectivity' declarations, e.g. '{V2E: v2e_table}'." + ) + return offset_provider[offset_tag] get_offset_type: Callable[[OffsetProviderType, str], OffsetProviderTypeElem] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and OffsetProviderType overlap @@ -1613,8 +1630,8 @@ def inverse_image(self, image_range: UnitRange | NamedRange) -> Sequence[NamedRa def premap( self, - index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], - *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + index_field: Connectivity | type[NeighborConnectivity], + *args: Connectivity | type[NeighborConnectivity], ) -> Connectivity: raise NotImplementedError() @@ -2056,7 +2073,7 @@ 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)] + return cls.bound_table()[cls._local()(int(item))] if "Local" in cls.__dict__: raise TypeError( f"'{cls.__qualname__}[{item!r}]': a connectivity is indexed by an integer" @@ -2073,23 +2090,29 @@ 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, tagged with `offset_tag`. + The type of the connectivity in DSL code: an offset from `Codomain` to `(Origin, Local)`. + + Its tag is `offset_tag`, which is how the IR names the connectivity and how the offset + provider is keyed once normalized (see `as_tag_keyed_offset_provider`). """ - from gt4py.next.ffront import fbuiltins - - 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()), - _derived=True, + from gt4py.next.type_system import type_specifications as ts + + local = cls._local() + return ts.OffsetType(source=cls.codomain, target=(cls.origin, local), tag=cls.offset_tag) + + def bound_table(cls) -> NeighborTable: + """The neighbor table bound to this connectivity in the current embedded execution.""" + from gt4py.next import embedded + + offset_provider = embedded.context.get_offset_provider(None) + if offset_provider is None: + raise RuntimeError( + f"'{cls.__qualname__}' can only be resolved to a table during embedded execution." ) - type.__setattr__(cls, "_field_offset", field_offset) - return field_offset + table = get_offset(offset_provider, cls.offset_tag) + assert is_neighbor_table(table) + return table class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( @@ -2293,3 +2316,128 @@ 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}" ) + + +@overload +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderLike, *, strict: bool = True +) -> OffsetProvider: ... +@overload +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderTypeLike, *, strict: bool = True +) -> OffsetProviderType: ... +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, *, strict: bool = True +) -> OffsetProvider | OffsetProviderType: + """ + Key an offset provider by tags, the form the IR and the backends use. + + A `NeighborConnectivity` key becomes its local dimension's tag. A string key is taken to be + such a tag already, and is rejected if it cannot be one: a tag is a qualified name, so a bare + name such as `"V2E"` is the removed `FieldOffset` spelling. + + Called on every program call, so it does not check tables against their declarations; see + `check_offset_provider`. + + Args: + offset_provider: The provider to normalize. + strict: Whether to reject string keys that cannot be tags. Internal hooks that are handed + hand-written providers, such as `embedded.context.update`, pass `False`. + """ + if not any(isinstance(key, ConnectivityMeta) for key in offset_provider): + if strict: + _check_tag_keys(offset_provider) + return offset_provider + result: dict[Tag, Any] = {} + for key, value in offset_provider.items(): + tag = key.offset_tag if isinstance(key, ConnectivityMeta) else key + if tag in result: + raise ValueError(f"The offset provider binds '{tag}' twice.") + result[tag] = value + if strict: + _check_tag_keys(result) + return result + + +def _check_tag_keys(offset_provider: Mapping[Any, Any]) -> None: + for key in offset_provider: + if not isinstance(key, str) or "." not in key: + raise TypeError( + f"Invalid offset-provider key '{key!r}': offset providers are keyed by" + " 'NeighborConnectivity' declarations, e.g. '{V2E: v2e_table}'. A bare name is the" + " spelling of the removed 'FieldOffset' (see ADR 0029)." + ) + + +def check_offset_provider(offset_provider: OffsetProviderLike | OffsetProviderTypeLike) -> None: + """ + Check every table of a tag-keyed offset provider against its connectivity declaration. + + A tag that does not name the local dimension of a declared connectivity -- e.g. one used only + by hand-written IR -- has no declaration to be checked against and is skipped. + + Raises: + ValueError: If a table does not match its declaration, see `check_neighbor_table`. + """ + for key, table in offset_provider.items(): + declaration: Any = key + if isinstance(key, str): + try: + declaration = _import_qualified_name(key) + except ValueError: + continue + if isinstance(declaration, DimensionMeta): + # the local dimension's tag names its owner's table + declaration = getattr(declaration, "owner", None) + if isinstance(declaration, ConnectivityMeta) and declaration.offset_tag in ( + key, + getattr(key, "offset_tag", None), + ): + check_neighbor_table(cast(type[NeighborConnectivity], declaration), table) + _check_shared_local_dimensions(offset_provider) + + +def _check_shared_local_dimensions( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, +) -> None: + """ + Check that the tables over one local dimension have the same neighbor structure. + + Reductions and sparse fields take the neighbor count and the skip values of a local dimension + from any one table over it (see `connectivity_key_over`), which is only sound if all of them + agree: the same number of neighbors, and a skip value at the same positions. + """ + by_local_dim: dict[Tag, list[tuple[Any, Any]]] = collections.defaultdict(list) + for key, table in offset_provider.items(): + if (neighbor_dim := _neighbor_dim_of(table)) is not None: + by_local_dim[neighbor_dim.tag].append((key, table)) + for local_tag, tables in by_local_dim.items(): + (first_key, first), *others = tables + first_type = first if isinstance(first, NeighborConnectivityType) else first.__gt_type__() + for key, table in others: + table_type = ( + table if isinstance(table, NeighborConnectivityType) else table.__gt_type__() + ) + same_structure = (table_type.max_neighbors, table_type.has_skip_values) == ( + first_type.max_neighbors, + first_type.has_skip_values, + ) + if ( + same_structure + and first_type.has_skip_values + and is_neighbor_table(first) + and is_neighbor_table(table) + ): + same_structure = bool( + np.array_equal( + first.asnumpy() == first_type.skip_value, + table.asnumpy() == table_type.skip_value, + ) + ) + if not same_structure: + raise ValueError( + f"'{key}' and '{first_key}' are bound to tables over the same local dimension" + f" '{local_tag}' with a different neighbor structure: connectivities sharing a" + " local dimension must have the same number of neighbors, and skip values at" + " the same positions." + ) diff --git a/src/gt4py/next/embedded/context.py b/src/gt4py/next/embedded/context.py index 8183a3292c..6ca2544ce9 100644 --- a/src/gt4py/next/embedded/context.py +++ b/src/gt4py/next/embedded/context.py @@ -73,7 +73,7 @@ def get_offset_provider(default: _T = _NO_DEFAULT_SENTINEL) -> common.OffsetProv def update( *, closure_column_range: common.NamedRange | eve.NothingType = eve.NOTHING, - offset_provider: common.OffsetProvider | eve.NothingType = eve.NOTHING, + offset_provider: common.OffsetProviderLike | eve.NothingType = eve.NOTHING, ) -> Generator[None, None, None]: """Context handler updating the current embedded context with the provided values.""" @@ -83,7 +83,9 @@ def update( closure_token = gtx_embedded.context._closure_column_range.set(closure_column_range) if offset_provider is not eve.NOTHING: assert not isinstance(offset_provider, eve.NothingType) - offset_provider_token = gtx_embedded.context._offset_provider.set(offset_provider) + offset_provider_token = gtx_embedded.context._offset_provider.set( + common.as_tag_keyed_offset_provider(offset_provider, strict=False) + ) try: yield None diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index 1bbdf3ef82..114442100f 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -239,9 +239,7 @@ def as_scalar(self) -> core_defs.ScalarT: def premap( self: NdArrayField, - *connectivities: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], + *connectivities: common.Connectivity | type[common.NeighborConnectivity], ) -> NdArrayField: """ Rearrange the field content using the provided connectivities (index mappings). @@ -316,13 +314,9 @@ def premap( codomains_counter: collections.Counter[common.Dimension] = collections.Counter() for connectivity in connectivities: - # For neighbor reductions, a FieldOffset or a connectivity declaration is passed - # instead of an actual Connectivity + # For neighbor reductions, a connectivity declaration is passed instead of a table 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() + connectivity = connectivity.bound_table() assert isinstance(connectivity, common.Connectivity) # Current implementation relies on skip_value == -1: @@ -371,10 +365,8 @@ def premap( def __call__( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: return functools.reduce( lambda field, current_index_field: field.premap(current_index_field), @@ -948,15 +940,12 @@ def _concat_where( NdArrayField.register_builtin_func(experimental.concat_where, _concat_where) # type: ignore[arg-type] -def _as_offset(offset: fbuiltins.FieldOffset, offset_field: NdArrayField) -> common.Connectivity: - if not fbuiltins.is_cartesian_offset(offset): - target_dims = ", ".join(d.__qualname__ for d in offset.target) # for the diagnostic - raise ValueError( - f"'as_offset' is only supported for Cartesian offsets " - f"(single target dimension equal to source dimension); " - f"got source '{offset.source.__qualname__}' and target ({target_dims})." - ) - source_dim = offset.source +def _as_offset(source_dim: common.Dimension, offset_field: NdArrayField) -> common.Connectivity: + if ( + not isinstance(source_dim, common.DimensionMeta) + or source_dim.kind is common.DimensionKind.LOCAL + ): + raise ValueError(f"'as_offset' shifts along a non-local dimension, got '{source_dim}'.") coords = _identity_index_array( offset_field.domain, source_dim, offset_field.array_ns, dtype=fbuiltins.IndexType ) diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index e9daedb8a9..37a14db1e1 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -160,9 +160,9 @@ def _make_compiled_programs_pool( def compile( self, - offset_provider: common.OffsetProviderType - | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + offset_provider: common.OffsetProviderTypeLike + | common.OffsetProviderLike + | list[common.OffsetProviderTypeLike | common.OffsetProviderLike] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -199,6 +199,7 @@ def compile( ) if not isinstance(offset_provider, list): offset_provider = [offset_provider] # type: ignore[list-item] # cleanup offset_provider vs offset_provider_type + offset_provider = [common.as_tag_keyed_offset_provider(op) for op in offset_provider] assert all( common.is_offset_provider(op) or common.is_offset_provider_type(op) @@ -376,12 +377,13 @@ def with_bound_args(self, **kwargs: Any) -> ProgramWithBoundArgs: def __call__( self, *args: Any, - offset_provider: common.OffsetProvider | None = None, + offset_provider: common.OffsetProviderLike | None = None, enable_jit: bool | None = None, **kwargs: Any, ) -> None: - if offset_provider is None: - offset_provider = {} + offset_provider = ( + {} if offset_provider is None else common.as_tag_keyed_offset_provider(offset_provider) + ) enable_jit = self.compilation_options.enable_jit if enable_jit is None else enable_jit with program_call_context( @@ -412,6 +414,7 @@ def __call__( stacklevel=2, ) + common.check_offset_provider(offset_provider) with next_embedded.context.update(offset_provider=offset_provider): with embedded_program_call_context(self, args, offset_provider, kwargs): self.definition_stage.definition(*args, **kwargs) @@ -431,7 +434,7 @@ class ProgramWithBoundArgs(Program): @override def __call__( - self, *args: Any, offset_provider: common.OffsetProvider | None = None, **kwargs: Any + self, *args: Any, offset_provider: common.OffsetProviderLike | None = None, **kwargs: Any ) -> None: if offset_provider is None: offset_provider = {} @@ -488,9 +491,9 @@ def __call__( @override def compile( self, - offset_provider: common.OffsetProviderType - | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + offset_provider: common.OffsetProviderTypeLike + | common.OffsetProviderLike + | list[common.OffsetProviderTypeLike | common.OffsetProviderLike] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -655,7 +658,9 @@ def __gt_closure_vars__(self) -> dict[str, Any]: def __call__(self, *args: Any, enable_jit: bool | None = None, **kwargs: Any) -> Any: if not next_embedded.context.within_valid_context() and self.backend is not None: # non embedded execution - offset_provider = {**kwargs.pop("offset_provider", {})} + offset_provider = { + **common.as_tag_keyed_offset_provider(kwargs.pop("offset_provider", {})) + } if "out" not in kwargs: raise errors.MissingArgumentError(None, "out", True) out = kwargs.pop("out") @@ -677,7 +682,10 @@ def __call__(self, *args: Any, enable_jit: bool | None = None, **kwargs: Any) -> else: if not next_embedded.context.within_valid_context(): # field_operator as program - kwargs["offset_provider"] = {**kwargs.pop("offset_provider", {})} + kwargs["offset_provider"] = { + **common.as_tag_keyed_offset_provider(kwargs.pop("offset_provider", {})) + } + common.check_offset_provider(kwargs["offset_provider"]) attributes = ( self.definition_stage.attributes if self.definition_stage diff --git a/src/gt4py/next/ffront/experimental.py b/src/gt4py/next/ffront/experimental.py index b547d67663..b1824b61f2 100644 --- a/src/gt4py/next/ffront/experimental.py +++ b/src/gt4py/next/ffront/experimental.py @@ -10,11 +10,16 @@ from gt4py._core import definitions as core_defs from gt4py.next import common, named_collections -from gt4py.next.ffront.fbuiltins import BuiltInFunction, FieldOffset, WhereBuiltinFunction +from gt4py.next.ffront.fbuiltins import BuiltInFunction, WhereBuiltinFunction @BuiltInFunction -def as_offset(offset: FieldOffset, field: common.Field, /) -> common.Connectivity: +def as_offset(dim: common.Dimension, field: common.Field, /) -> common.Connectivity: + """ + Shift along `dim` by the per-point amounts in the integer `field`. + + `a(as_offset(KDim, k_offsets))` reads `a` at `k + k_offsets[k]` in `KDim`. + """ raise NotImplementedError() diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 0916417a5f..cf7d4a2ff4 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -7,11 +7,9 @@ # SPDX-License-Identifier: BSD-3-Clause import dataclasses -import functools import inspect import math import operator -import warnings from builtins import bool, float, int, tuple # noqa: A004 shadowing a Python built-in from types import UnionType from typing import ( @@ -36,7 +34,6 @@ from gt4py._core import definitions as core_defs from gt4py.next import common, named_collections from gt4py.next.common import Dimension, Field # noqa: F401 [unused-import] for TYPE_BUILTINS -from gt4py.next.iterator import runtime from gt4py.next.type_system import type_specifications as ts @@ -130,8 +127,6 @@ def _type_conversion_helper(t: type) -> type[ts.TypeSpec] | tuple[type[ts.TypeSp return ts.FieldType elif t is common.Dimension: return ts.DimensionType - elif t is FieldOffset: - return ts.OffsetType elif t is common.Connectivity: return ts.OffsetType elif t is core_defs.ScalarT: @@ -473,91 +468,3 @@ def impl( assert (diff := actual_export - should_export) == set(), ( f"Symbol(s) exported but not defined in 'fbuiltins': {diff}" ) - - -# TODO(tehrengruber): FieldOffset and runtime.Offset are not an exact conceptual -# match. Revisit if we want to continue subclassing here. If we split -# them also check whether Dimension should continue to be the shared or define -# guidelines for decision. -@dataclasses.dataclass(frozen=True) -class FieldOffset(runtime.Offset): - #: The tag, i.e. the offset-provider key. - value: str - source: common.Dimension - target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] - #: Set when derived from a `NeighborConnectivity` declaration, which is not deprecated. - _derived: bool = dataclasses.field(default=False, repr=False, compare=False, kw_only=True) - - @functools.cached_property - def _cache(self) -> dict: - return {} - - def __post_init__(self) -> None: - if len(self.target) == 2: - if self.target[1].kind != common.DimensionKind.LOCAL: - raise ValueError("Second dimension in offset must be a local dimension.") - if not self._derived: - warnings.warn( - "Declaring an unstructured connectivity with 'FieldOffset' is deprecated;" - " declare a 'NeighborConnectivity' class instead (see ADR 0029):\n" - " class V2E(gtx.NeighborConnectivity[Vertex, Edge]):\n" - " class Local(gtx.LocalDimensionIndex): ...", - DeprecationWarning, - stacklevel=3, - ) - - 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 - - assert isinstance(self.value, str) - current_offset_provider = embedded.context.get_offset_provider(None) - assert current_offset_provider is not None - offset_definition = common.get_offset(current_offset_provider, self.value) - - assert common.is_neighbor_table(offset_definition) - named_index = (self.target[-1])(offset) - connectivity = offset_definition[named_index] - - return connectivity - - def as_connectivity_field(self) -> common.Connectivity: - """Convert to connectivity field using the offset providers in current embedded execution context.""" - from gt4py.next import embedded # avoid circular import - - assert isinstance(self.value, str) - current_offset_provider = embedded.context.get_offset_provider(None) - assert current_offset_provider is not None - offset_definition = common.get_offset(current_offset_provider, self.value) - - cache_key = id(offset_definition) - if (connectivity := self._cache.get(cache_key, None)) is None: - if isinstance(offset_definition, common.Connectivity): - connectivity = offset_definition - else: - raise NotImplementedError() - - self._cache[cache_key] = connectivity - - return connectivity - - -def is_cartesian_offset(offset: FieldOffset | ts.OffsetType) -> bool: - 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 - ) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 85eaafe610..bbdf8433fe 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -997,16 +997,14 @@ 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_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 + if not isinstance(arg_0, ts.DimensionType) or arg_0.dim.kind is common.DimensionKind.LOCAL: 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"'as_offset' shifts along a non-local dimension, e.g. 'as_offset(KDim, field)';" + f" got '{arg_0}'.", ) + dim = arg_0.dim if not type_info.is_integral(arg_1): raise errors.DSLError( node.location, @@ -1015,16 +1013,20 @@ 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 dim 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"'{dim}' not in list of offset field dimensions '{arg_1.dims}'. " f"{node.location}", ) return foast.Call( - func=node.func, args=node.args, kwargs=node.kwargs, type=arg_0, location=node.location + func=node.func, + args=node.args, + kwargs=node.kwargs, + type=ts.OffsetType(source=dim, target=(dim,)), + location=node.location, ) def _deduce_where_return_type( diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 841ce4ca98..1b5aabf37d 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -331,9 +331,9 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: # `field(as_offset(Off, offset_field))` 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 + dim_type = func_args.args[0].type + assert isinstance(dim_type, ts.DimensionType) + dim = dim_type.dim offset_field = self.visit(func_args.args[1], **kwargs) current_expr = im.as_fieldop( im.lambda_("__it", "__offset")( diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 131647baf4..a00eccf544 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -17,7 +17,6 @@ from gt4py.eve import NodeTranslator, traits from gt4py.next import common, config, errors, utils from gt4py.next.ffront import ( - fbuiltins, gtcallable, program_ast as past, stages as ffront_stages, @@ -74,7 +73,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.ConnectivityMeta, common.DimensionMeta + all_closure_vars, 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 4e24d83881..26bfc2f765 100644 --- a/src/gt4py/next/ffront/transform_utils.py +++ b/src/gt4py/next/ffront/transform_utils.py @@ -10,7 +10,6 @@ from typing import Any, Iterable, Optional from gt4py.next import common -from gt4py.next.ffront import fbuiltins from gt4py.next.ffront.gtcallable import GTCallable @@ -47,9 +46,7 @@ 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 | type[common.NeighborConnectivity] | common.Dimension - ], + offsets_and_dimensions: Iterable[type[common.NeighborConnectivity] | common.Dimension], ) -> common.GridType: """ Derive grid type from actually occurring dimensions and check against optional user request. @@ -61,9 +58,7 @@ def _deduce_grid_type( deduced_grid_type = common.GridType.CARTESIAN for o in offsets_and_dimensions: - if isinstance(o, common.ConnectivityMeta) or ( - isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o) - ): + if isinstance(o, common.ConnectivityMeta): deduced_grid_type = common.GridType.UNSTRUCTURED break if isinstance(o, common.DimensionMeta) and o.kind == common.DimensionKind.LOCAL: @@ -75,7 +70,7 @@ def _deduce_grid_type( and deduced_grid_type == common.GridType.UNSTRUCTURED ): raise ValueError( - "'grid_type == GridType.CARTESIAN' was requested, but unstructured 'FieldOffset' or local 'Dimension' was found." + "'grid_type == GridType.CARTESIAN' was requested, but a 'NeighborConnectivity' or a local dimension was found." ) return deduced_grid_type if requested_grid_type is None else requested_grid_type diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 7062053223..dced958782 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -51,7 +51,6 @@ exceptions as embedded_exceptions, operators, ) -from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import builtins, runtime from gt4py.next.type_system import type_specifications as ts, type_translation @@ -156,9 +155,7 @@ def asnumpy(self) -> np.ndarray: def premap( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError @@ -176,10 +173,8 @@ def as_scalar(self) -> xtyping.Never: def __call__( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError() @@ -1165,10 +1160,8 @@ def as_scalar(self) -> core_defs.IntegralScalar: def premap( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1308,10 +1301,8 @@ def asnumpy(self) -> np.ndarray: def premap( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1475,7 +1466,9 @@ def _as_offset_tag( @builtins.neighbors.register(EMBEDDED) def neighbors(offset: runtime.Offset | type[common.NeighborConnectivity], it: ItIterator) -> _List: field_offset: runtime.Offset = ( - offset.__gt_field_offset__() if isinstance(offset, common.ConnectivityMeta) else offset + runtime.Offset(value=offset.offset_tag) + if isinstance(offset, common.ConnectivityMeta) + else offset ) offset_str = _as_offset_tag(field_offset) assert isinstance(offset_str, str) diff --git a/src/gt4py/next/iterator/runtime.py b/src/gt4py/next/iterator/runtime.py index 88a466229f..9df3f64e1a 100644 --- a/src/gt4py/next/iterator/runtime.py +++ b/src/gt4py/next/iterator/runtime.py @@ -78,7 +78,10 @@ def __call__( offset_provider=None, column_axis=None, ): - offset_provider = offset_provider or self.offset_provider + # NOTE: not strict: iterator IR names offsets by arbitrary strings. + offset_provider = common.as_tag_keyed_offset_provider( + offset_provider or self.offset_provider, strict=False + ) column_axis = column_axis or self.column_axis if backend is not None: diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index 0b49f9ec0c..f5bb207efb 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -643,6 +643,7 @@ def _compile_variant( else: raise ValueError(f"Invalid 'offset_provider': {offset_provider}") + common.check_offset_provider(offset_provider) self._initialize_argument_descriptor_mapping(argument_descriptors) _validate_argument_descriptors(self.program_type, argument_descriptors) diff --git a/src/gt4py/next/otf/options.py b/src/gt4py/next/otf/options.py index 4f77d44586..c368e524f7 100644 --- a/src/gt4py/next/otf/options.py +++ b/src/gt4py/next/otf/options.py @@ -15,7 +15,7 @@ class CompilationOptionsArgs(TypedDict, total=False): enable_jit: bool static_params: Sequence[str] - connectivities: common.OffsetProvider + connectivities: common.OffsetProviderLike static_domains: bool @@ -36,9 +36,15 @@ class CompilationOptions: #: 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-. - connectivities: common.OffsetProvider | None = None + connectivities: common.OffsetProviderLike | None = None static_domains: bool = False + def __post_init__(self) -> None: + if self.connectivities is not None: + object.__setattr__( + self, "connectivities", common.as_tag_keyed_offset_provider(self.connectivities) + ) + assert CompilationOptionsArgs.__annotations__.keys() == CompilationOptions.__annotations__.keys() 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 36a80cb97c..c4560a60fe 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -17,7 +17,6 @@ from gt4py._core import definitions as core_defs from gt4py.eve import codegen from gt4py.next import common -from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import ir as itir from gt4py.next.iterator.transforms import pass_manager from gt4py.next.otf import artifacts, stages, workflow @@ -81,21 +80,14 @@ def _process_regular_arguments( if isinstance(parameter.type_, ts.FieldType): for dim in parameter.type_.dims: - if ( - isinstance( - dim, fbuiltins.FieldOffset - ) # TODO(havogt): remove support for FieldOffset as Dimension - or dim.kind is common.DimensionKind.LOCAL - ): + if dim.kind is common.DimensionKind.LOCAL: # translate sparse dimensions to tuple dtype - # 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 + # NOTE: the local dimension's tag names the `generated::_t` tag type + # (mangled); its table may be keyed by a connectivity sharing it. + dim_name = dim.tag connectivity = common.get_offset_type( offset_provider_type, - dim_name - if isinstance(dim, fbuiltins.FieldOffset) - else common.connectivity_key_over(offset_provider_type, dim), + common.connectivity_key_over(offset_provider_type, dim), ) assert isinstance(connectivity, common.NeighborConnectivityType) size = connectivity.max_neighbors diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 2efc73b47c..3f63cec788 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -62,7 +62,6 @@ "Vertex", "Edge", "Cell", - "EdgeOffset", "MeshDescriptor", "CartesianGridDescriptor", ] @@ -187,12 +186,6 @@ class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... KHalfDim = common.flip_staggered(KDim) -Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) -Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) - - -EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) - class C2V(gtx.NeighborConnectivity[Cell, Vertex]): class Local(gtx.LocalDimensionIndex): ... @@ -230,7 +223,10 @@ def simple_cartesian_grid( name="simple_cartesian_grid", sizes=sizes, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + offset_provider_type=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -316,28 +312,28 @@ def simple_mesh(allocator) -> MeshDescriptor: e2v_arr = np.asarray(e2v_arr, dtype=gtx.IndexType) offset_provider = { - V2E.Local.tag: constructors.as_connectivity( + V2E: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 4}, codomain=Edge, data=v2e_arr, skip_value=None, allocator=allocator, ), - E2V.Local.tag: constructors.as_connectivity( + E2V: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.Local.tag: constructors.as_connectivity( + C2V: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 4}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.Local.tag: constructors.as_connectivity( + C2E: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 4}, codomain=Edge, data=c2e_arr, @@ -352,7 +348,10 @@ def simple_mesh(allocator) -> MeshDescriptor: num_edges=np.int32(num_edges), num_cells=num_cells, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + offset_provider_type=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -411,28 +410,28 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) offset_provider = { - V2E.Local.tag: constructors.as_connectivity( + V2E: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 5}, codomain=Edge, data=v2e_arr, skip_value=common._DEFAULT_SKIP_VALUE, allocator=allocator, ), - E2V.Local.tag: constructors.as_connectivity( + E2V: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.Local.tag: constructors.as_connectivity( + C2V: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 3}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.Local.tag: constructors.as_connectivity( + C2E: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 3}, codomain=Edge, data=c2e_arr, @@ -447,7 +446,10 @@ def skip_value_mesh(allocator) -> MeshDescriptor: num_edges=num_edges, num_cells=num_cells, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + offset_provider_type=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -460,3 +462,18 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) def mesh_descriptor(request, exec_alloc_descriptor) -> MeshDescriptor: yield request.param(exec_alloc_descriptor.allocator) + + +def ir_level(mesh: MeshDescriptor) -> MeshDescriptor: + """ + A copy of `mesh` whose offset provider is keyed by tags, as the IR-level APIs expect. + + User-facing entry points normalize a class-keyed provider themselves; tests that drive the + lowering or the backends directly have to hand them the tag-keyed form. + """ + return types.SimpleNamespace( + **{ + **vars(mesh), + "offset_provider": common.as_tag_keyed_offset_provider(mesh.offset_provider), + } + ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py index 1e65cd734c..c156a61407 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py @@ -19,8 +19,6 @@ cartesian_case, ) from next_tests.integration_tests.cases_utils import ( - Ioff, - Koff, exec_alloc_descriptor, ) @@ -62,10 +60,10 @@ def test_offset_field(cartesian_case): @gtx.field_operator def testee(a: cases.IKField, offset_field: cases.IKField) -> gtx.Field[[IDim, KDim], bool]: - a_i = a(as_offset(Ioff, offset_field)) + a_i = a(as_offset(IDim, offset_field)) # note: this leads to an access to offset_field in # IDim: (0, out.size[I]), KDim: (0, out.size[K]+1) - a_i_k = a_i(as_offset(Koff, offset_field)) + a_i_k = a_i(as_offset(KDim, offset_field)) b_i = a(IDim + 1) b_i_k = b_i(KDim + 1) return a_i_k == b_i_k @@ -97,7 +95,7 @@ def test_offset_field_of_chained_ops(cartesian_case): def testee(a: cases.IKField, offset_field: cases.IKField) -> cases.IKField: b = a + 1 c = b * 2 - return c(as_offset(Koff, offset_field)) + return c(as_offset(KDim, offset_field)) out = cases.allocate(cartesian_case, testee, cases.RETURN)() a = cases.allocate(cartesian_case, testee, "a").extend({KDim: (0, 1)})() diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py index 57a2af8dfc..68c2dcf72b 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py @@ -252,7 +252,7 @@ def test_compile_unstructured(unstructured_case, compile_testee_unstructured): compile_testee_unstructured(*args, offset_provider=unstructured_case.offset_provider, **kwargs) - v2e_numpy = unstructured_case.offset_provider[V2E.Local.tag].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), @@ -317,7 +317,7 @@ def test_compile_unstructured_for_two_offset_providers( *args, offset_provider=unstructured_case.offset_provider, **kwargs ) - v2e_numpy = unstructured_case.offset_provider[V2E.Local.tag].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py index dfe013e45a..365de7d473 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py @@ -454,7 +454,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: t = concat_where(Vertex < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -477,7 +477,7 @@ def testee( t = concat_where(KDim < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() k_mask = np.arange(unstructured_case_3d.default_sizes[KDim]) < 2 cases.verify_with_default_data( unstructured_case_3d, @@ -497,7 +497,7 @@ def test_with_local_and_nonlocal_field(unstructured_case, static_domains: bool): def testee(a: cases.EField, b: cases.VField) -> cases.VField: return neighbor_sum(concat_where(Vertex < 2, a(V2E), b), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -524,7 +524,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), b(V2E)), (c(V2E), d(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -557,7 +557,7 @@ def testee(a: cases.EField) -> tuple[cases.VField, cases.VField]: neighbor_sum(concat_where(Vertex < 2, 3, a(V2E)), axis=V2EDim), ) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -590,7 +590,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), c), (3, b(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py index 753015f0d7..b65a16ffea 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py @@ -31,11 +31,11 @@ def testee( ) # multiplication with shifted `ones` because reduction of only non-shifted field with local dimension is not supported inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) ones = cases.allocate(unstructured_case, testee, "ones").strategy(cases.ConstInitializer(1))() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify( unstructured_case, testee, @@ -55,7 +55,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return inp[V2EDim(0)] + inp[V2EDim(1)] + inp[V2EDim(2)] + inp[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -77,7 +77,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int64 return inp_64[V2EDim(0)] + inp_64[V2EDim(1)] + inp_64[V2EDim(2)] + inp_64[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -99,7 +99,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return neighbor_sum(inp, axis=V2EDim) inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -107,7 +107,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 testee, inp, out=cases.allocate(unstructured_case, testee, cases.RETURN)(), - ref=np.sum(unstructured_case.offset_provider[V2EDim.tag].asnumpy(), axis=1), + ref=np.sum(unstructured_case.offset_provider[V2E].asnumpy(), axis=1), ) @@ -119,7 +119,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: return inp(V2E) out = unstructured_case.as_field( - [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider[V2EDim.tag].asnumpy()) + [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider[V2E].asnumpy()) ) inp = cases.allocate(unstructured_case, testee, "inp")() cases.verify( @@ -127,5 +127,5 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: testee, inp, out=out, - ref=inp.asnumpy()[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], + ref=inp.asnumpy()[unstructured_case.offset_provider[V2E].asnumpy()], ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py index 262751f234..dbc623f2de 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py @@ -60,7 +60,7 @@ def test_import_offset_module_unstructured_shift(unstructured_case): def testee(a: cases.EField) -> cases.VField: return neighbor_sum(a(cases.V2E), axis=cases.V2EDim) - v2e_table = unstructured_case.offset_provider[cases.V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[cases.V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -77,7 +77,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider[cases.E2VDim.tag].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[cases.E2V].asnumpy()[:, 0]], ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py index 39aee5de06..6ceead30d6 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py @@ -483,7 +483,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) ) @@ -641,7 +641,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, 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 bce54f4e6a..07793c44fe 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 @@ -42,7 +42,7 @@ class V2EShared(gtx.NeighborConnectivity[V, E]): @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() + v2e_arr = mesh.offset_provider[cases_utils.V2E].asnumpy() table = constructors.as_connectivity( domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, codomain=E, @@ -65,9 +65,7 @@ def case(exec_alloc_descriptor): 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.offset_tag: table, V2EShared.offset_tag: shared_table}, + offset_provider={V2E: table, V2EShared: 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, @@ -75,7 +73,7 @@ def case(exec_alloc_descriptor): def _table(case: cases.Case, connectivity=V2E) -> np.ndarray: - return case.offset_provider[connectivity.offset_tag].asnumpy() + return case.offset_provider[connectivity].asnumpy() @pytest.mark.uses_unstructured_shift @@ -148,9 +146,7 @@ def testee( @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]} - ) + return dataclasses.replace(case, offset_provider={V2EShared: case.offset_provider[V2EShared]}) @pytest.mark.uses_unstructured_shift diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py index dd0cc6fb43..f0ce221851 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py @@ -53,7 +53,7 @@ def testee(edge_f: cases.EField) -> cases.VField: inp = cases.allocate(unstructured_case, testee, "edge_f", strategy=strategy)() out = cases.allocate(unstructured_case, testee, cases.RETURN)() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() ref = np.max( inp.asnumpy()[v2e_table], axis=1, @@ -70,7 +70,7 @@ def minover(edge_f: cases.EField) -> cases.VField: out = min_over(edge_f(V2E), axis=V2EDim) return out - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, minover, @@ -100,7 +100,7 @@ def reduction_ek_field( "fop", [reduction_e_field, reduction_ek_field], ids=lambda fop: fop.__name__ ) def test_neighbor_sum(unstructured_case_3d, fop): - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() edge_f = cases.allocate(unstructured_case_3d, fop, "edge_f")() @@ -152,7 +152,7 @@ def fencil_op(edge_f: EKField) -> VKField: def fencil(edge_f: EKField, out: VKField): fencil_op(edge_f, out=out) - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() field = cases.allocate(unstructured_case_3d, fencil, "edge_f", sizes={KDim: 2})() out = cases.allocate(unstructured_case_3d, fencil_op, cases.RETURN, sizes={KDim: 1})() @@ -185,7 +185,7 @@ def reduce_expr(edge_f: cases.EField) -> cases.VField: def fencil(edge_f: cases.EField, out: cases.VField): reduce_expr(edge_f, out=out) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, fencil, @@ -207,7 +207,7 @@ def test_reduction_with_common_expression(unstructured_case): def testee(flux: cases.EField) -> cases.VField: return neighbor_sum(flux(V2E) + flux(V2E), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -223,7 +223,7 @@ def test_reduction_expression_with_where(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, inp(V2E), inp(V2E)), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -252,7 +252,7 @@ def test_reduction_expression_with_where_and_tuples(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, (inp(V2E), inp(V2E)), (inp(V2E), inp(V2E)))[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -281,7 +281,7 @@ def test_reduction_expression_with_where_and_scalar(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(inp(V2E) + where(mask, inp(V2E), 1), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -326,7 +326,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]], ) @@ -347,7 +347,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() neighbor_0_iter = iter(enumerate(e2v_table[:, 0])) edge_start = next(i for i, v in neighbor_0_iter if v >= ORIGIN) edge_stop = next(i for i, v in neighbor_0_iter if v < ORIGIN) @@ -392,16 +392,16 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_flat, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], ) cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_intermediate_result, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], comparison=lambda inp, tmp: np.all(inp == tmp), ) @@ -409,8 +409,8 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], ) @@ -432,7 +432,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() neighbor_iter = iter(enumerate(e2v_table)) edge_start = next(i for i, v in neighbor_iter if all(v >= ORIGIN)) edge_stop = next(i for i, v in neighbor_iter if any(v < ORIGIN)) @@ -453,12 +453,11 @@ def testee(a: cases.VField) -> cases.VField: unstructured_case, testee, ref=lambda a: np.sum( - np.sum(a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()], axis=1, initial=0)[ - unstructured_case.offset_provider[V2EDim.tag].asnumpy() + np.sum(a[unstructured_case.offset_provider[E2V].asnumpy()], axis=1, initial=0)[ + unstructured_case.offset_provider[V2E].asnumpy() ], axis=1, - where=unstructured_case.offset_provider[V2EDim.tag].asnumpy() - != common._DEFAULT_SKIP_VALUE, + where=unstructured_case.offset_provider[V2E].asnumpy() != common._DEFAULT_SKIP_VALUE, ), comparison=lambda a, tmp_2: np.all(a == tmp_2), ) @@ -479,8 +478,8 @@ def testee(inp: cases.EField) -> cases.EField: unstructured_case, testee, ref=lambda inp: np.sum( - np.sum(inp[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1)[ - unstructured_case.offset_provider[E2VDim.tag].asnumpy() + np.sum(inp[unstructured_case.offset_provider[V2E].asnumpy()], axis=1)[ + unstructured_case.offset_provider[E2V].asnumpy() ], axis=1, ), @@ -497,7 +496,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: tmp = red(E2V[0]) return tmp - v2e = unstructured_case.offset_provider[V2EDim.tag] + v2e = unstructured_case.offset_provider[V2E] cases.verify_with_default_data( unstructured_case, reduce_tuple_element, @@ -506,7 +505,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: axis=1, initial=0, where=v2e.asnumpy() != common._DEFAULT_SKIP_VALUE, - )[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], + )[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]], ) @@ -518,7 +517,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: tmp = neighbor_sum(b(V2E) if 2 < 3 else a(V2E), axis=V2EDim) return tmp - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -540,7 +539,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex], int32]: inp = cases.allocate(unstructured_case, testee, "inp")() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify( unstructured_case, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py index ac348266f3..15a94f2e36 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py @@ -26,7 +26,7 @@ unstructured_case, unstructured_case_3d, ) -from next_tests.integration_tests.cases_utils import Koff, exec_alloc_descriptor, mesh_descriptor +from next_tests.integration_tests.cases_utils import exec_alloc_descriptor, mesh_descriptor @pytest.mark.uses_cartesian_shift @@ -156,7 +156,7 @@ def test_cartesian_half_shift_as_offset(cartesian_case): def testee( a: gtx.Field[[IDim, KHalfDim], np.int32], offset_field: cases.IKField ) -> cases.IKField: - return a(KDim - 0.5)(as_offset(Koff, offset_field)) + return a(KDim - 0.5)(as_offset(KDim, offset_field)) ksize = cartesian_case.default_sizes[KDim] a = cases.allocate(cartesian_case, testee, "a", sizes={KHalfDim: ksize + 1})() @@ -180,7 +180,7 @@ def testee( a: gtx.Field[[IDim, KHalfDim], np.int32], offset_field: cases.IKField ) -> cases.IKField: b = a + 1 - return b(KDim - 0.5)(as_offset(Koff, offset_field)) + return b(KDim - 0.5)(as_offset(KDim, offset_field)) ksize = cartesian_case.default_sizes[KDim] a = cases.allocate(cartesian_case, testee, "a", sizes={KHalfDim: ksize + 1})() @@ -208,7 +208,7 @@ def testee( a: gtx.Field[[Vertex, KHalfDim], np.int32], offset_field: gtx.Field[[Edge, KDim], np.int32], ) -> gtx.Field[[Edge, KDim], np.int32]: - return a(E2V[0])(KDim - 0.5)(as_offset(Koff, offset_field)) + return a(E2V[0])(KDim - 0.5)(as_offset(KDim, offset_field)) nvertices = unstructured_case_3d.default_sizes[Vertex] ksize = unstructured_case_3d.default_sizes[KDim] @@ -223,7 +223,7 @@ def testee( )() out = cases.allocate(unstructured_case_3d, testee, cases.RETURN)() - e2v_table = unstructured_case_3d.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case_3d.offset_provider[E2V].asnumpy() cases.verify( unstructured_case_3d, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py index 0eb876abde..20ef5e789a 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py @@ -83,9 +83,7 @@ def test_verification(testee, exec_alloc_descriptor, mesh_descriptor): a = cases.allocate(unstructured_case, testee, "a")() out = cases.allocate(unstructured_case, testee, "out")() - first_nbs, second_nbs = ( - mesh_descriptor.offset_provider[E2VDim.tag].asnumpy()[:, i] for i in [0, 1] - ) + first_nbs, second_nbs = (mesh_descriptor.offset_provider[E2V].asnumpy()[:, i] for i in [0, 1]) ref = (a.ndarray * 2)[first_nbs] + (a.ndarray * 2)[second_nbs] cases.verify( @@ -105,7 +103,7 @@ def test_temporary_symbols(testee, mesh_descriptor): gtir_with_tmp = apply_common_transforms( testee.gtir, extract_temporaries=True, - offset_provider=mesh_descriptor.offset_provider, + offset_provider=common.as_tag_keyed_offset_provider(mesh_descriptor.offset_provider), ) params = ["num_vertices", "num_edges", "num_cells"] diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index c55a145314..ed91a34e9e 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -164,8 +164,8 @@ def testee(a: cases.EField, b: cases.EField) -> tuple[cases.VField, cases.VField unstructured_case, testee, ref=lambda a, b: [ - np.sum(a[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), - np.sum(b[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), + np.sum(a[unstructured_case.offset_provider[V2E].asnumpy()], axis=1), + np.sum(b[unstructured_case.offset_provider[V2E].asnumpy()], axis=1), ], comparison=lambda a, tmp: (np.all(a[0] == tmp[0]), np.all(a[1] == tmp[1])), ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py index 6a160cefc3..2e4f3d8b78 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py @@ -51,7 +51,7 @@ def testee(a: gtx.Field[[Vertex], np.float64]) -> gtx.Field[[Edge], int64]: tmp = astype(a(E2V), int64) return neighbor_sum(tmp, axis=E2VDim) - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py index 01e4ed92bb..eabf67327d 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py @@ -546,7 +546,7 @@ def test_program_unstructured(unstructured_case): unstructured_case.default_sizes[Cell], unstructured_case.default_sizes[Edge], inout=(out_a_shifted, out_a), - ref=((a.ndarray)[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]], a), + ref=((a.ndarray)[unstructured_case.offset_provider[C2E].asnumpy()[:, 1]], a), ) @@ -600,7 +600,7 @@ def test_program_temporary(unstructured_case): extend={Cell: (-restrict_cell[0], restrict_cell[1])}, )() - e2v = (a.ndarray)[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 1]] + e2v = (a.ndarray)[unstructured_case.offset_provider[E2V].asnumpy()[:, 1]] cases.verify( unstructured_case, prog_temporary, @@ -616,7 +616,7 @@ def test_program_temporary(unstructured_case): inout=(out_edge, out_cell), ref=( e2v[restrict_edge[0] : edge_size + restrict_edge[1]], - e2v[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]][ + e2v[unstructured_case.offset_provider[C2E].asnumpy()[:, 1]][ restrict_cell[0] : cell_size + restrict_cell[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 df066e21cf..a33f4d0f12 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 @@ -7,19 +7,13 @@ # SPDX-License-Identifier: BSD-3-Clause """ -Regression tests for the four independently authored names of one connectivity. +Regression tests for the names under which one connectivity is used. -Using a single connectivity requires four strings to agree, none of which is -checked against the others at declaration time: - - N1 the `FieldOffset` tag `FieldOffset("V2E", ...)` - N2 the Python variable it is bound to `V2E = FieldOffset(...)` - N3 the local dimension's name `Dimension("V2E", kind=LOCAL)` - N4 the offset-provider key `offset_provider={"V2E": ...}` - -The `V2EDim = Dimension("V2E")` convention makes all four equal, which hides -which one each execution path actually uses. These tests break the convention -deliberately, one name at a time, so the real requirement is visible. +With `FieldOffset`, using a connectivity required four independently authored strings to +agree -- the offset tag, the Python variable it was bound to, the local dimension's name and +the offset-provider key -- and each execution path silently depended on a different subset of +them. A `NeighborConnectivity` declaration produces all of them (ADR 0029), so what is left to +pin is that the *Python* name a declaration is reached through does not matter. """ import numpy as np @@ -42,26 +36,22 @@ class V(gtx.DimensionIndex): ... class E(gtx.DimensionIndex): ... -#: N1 == N3 == N4, but N2 differs: the tag is `TaggedOffDim.tag`, the variable is `off_a`. -class TaggedOffDim(gtx.LocalDimensionIndex): ... - - -off_a = gtx.FieldOffset(TaggedOffDim.tag, source=E, target=(V, TaggedOffDim)) +class V2E(gtx.NeighborConnectivity[V, E]): + class Local(gtx.LocalDimensionIndex): ... -#: N1 == N2 == N4, but N3 differs: the local dimension is `Neigh`, the tag is `OffB`. -class Neigh(gtx.LocalDimensionIndex): ... +#: The declaration, reached through a different Python name. +off_a = V2E +#: Its local dimension, likewise. +Neigh = V2E.Local -OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) - - -def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimension) -> cases.Case: - """A `Case` binding the same table under each of `tags`.""" +@pytest.fixture +def case(exec_alloc_descriptor) -> cases.Case: 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. - v2e_arr = mesh.offset_provider[cases_utils.V2EDim.tag].asnumpy() + v2e_arr = mesh.offset_provider[cases_utils.V2E].asnumpy() return cases.Case( ( None @@ -69,14 +59,13 @@ def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimens else exec_alloc_descriptor ), offset_provider={ - tag: constructors.as_connectivity( - domain={V: v2e_arr.shape[0], local_dim: v2e_arr.shape[1]}, + off_a: constructors.as_connectivity( + domain={V: v2e_arr.shape[0], Neigh: v2e_arr.shape[1]}, codomain=E, data=v2e_arr, 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, @@ -84,82 +73,29 @@ def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimens ) -@pytest.fixture -def case_tag_vs_variable_name(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, (TaggedOffDim.tag,), TaggedOffDim) - - -@pytest.fixture -def case_tag_vs_local_dim(exec_alloc_descriptor): - # 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: - return case.offset_provider[tag].asnumpy() - +def _neighbor_table(case: cases.Case) -> np.ndarray: + return case.offset_provider[V2E].asnumpy() -# --- N2: the tag differs from the Python variable name ---------------------------- -# Lowering used to emit the *variable* name as the IR shift tag, so embedded and -# compiled execution of the same program needed different provider keys. - -def test_shift_tag_differs_from_variable_name(case_tag_vs_variable_name): +def test_shift_through_an_alias(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: return a(off_a[1]) - cases.verify_with_default_data( - case_tag_vs_variable_name, - foo, - lambda a: a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)[:, 1]], - ) - - -def test_reduction_tag_differs_from_variable_name(case_tag_vs_variable_name): - @gtx.field_operator - def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return neighbor_sum(a(off_a), axis=TaggedOffDim) - - cases.verify_with_default_data( - case_tag_vs_variable_name, - foo, - lambda a: np.sum(a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)], axis=1), - ) - - -# --- N3: the tag differs from the local dimension's name -------------------------- -# The shape of a connectivity sharing another one's local dimension. - + cases.verify_with_default_data(case, foo, lambda a: a[_neighbor_table(case)[:, 1]]) -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, - gtfn would silently ignore the neighbor index, see - https://github.com/GridTools/gridtools/pull/1814. - """ +def test_reduction_through_an_alias(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return a(OffB[1]) + return neighbor_sum(a(off_a), axis=Neigh) - cases.verify_with_default_data( - case_tag_vs_local_dim, - foo, - lambda a: a[_neighbor_table(case_tag_vs_local_dim, "OffB")[:, 1]], - ) + cases.verify_with_default_data(case, foo, lambda a: np.sum(a[_neighbor_table(case)], axis=1)) -def test_reduction_tag_differs_from_local_dim_name(case_tag_vs_local_dim): +def test_reduction_over_the_nested_name(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return neighbor_sum(a(OffB), axis=Neigh) + return neighbor_sum(a(V2E), axis=off_a.Local) - cases.verify_with_default_data( - case_tag_vs_local_dim, - foo, - lambda a: np.sum(a[_neighbor_table(case_tag_vs_local_dim, "OffB")], axis=1), - ) + cases.verify_with_default_data(case, foo, lambda a: np.sum(a[_neighbor_table(case)], axis=1)) 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 3acc165129..1fa34616de 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 @@ -827,7 +827,6 @@ def test_premap_disjoint_inverse_image_raises(): def test_as_offset_1d(): # Dynamic per-point shift along I: out[i] == f[i + off[i]], full domain when all shifts in-bounds. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -835,7 +834,7 @@ def test_as_offset_1d(): off_arr = np.asarray([1, 0, -1, 0, 1, 0, -1, 0, 1, 0], dtype=int) off = common._field(off_arr, domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),))) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) assert np.all(result.ndarray == f.ndarray[np.arange(10) + off_arr]) @@ -843,7 +842,6 @@ def test_as_offset_1d(): def test_as_offset_narrow_offset_dtype_no_wrap(): # An int8 offset field over a domain larger than 128 must not wrap into the index table. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) N = 200 f = common._field( @@ -852,7 +850,7 @@ def test_as_offset_narrow_offset_dtype_no_wrap(): off_arr = np.zeros(N, dtype=np.int8) off = common._field(off_arr, domain=common.Domain(dims=(I,), ranges=(UnitRange(0, N),))) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(0, N),)) assert np.all(result.ndarray == f.ndarray) @@ -860,7 +858,6 @@ def test_as_offset_narrow_offset_dtype_no_wrap(): def test_as_offset_2d_shift_one_keep_other(): # Shift along I by a per-(i, j) offset, leave J: out[i, j] == f[i + off[i, j], j]. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) NI, NJ = 4, 3 dom = common.Domain(dims=(I, J), ranges=(UnitRange(0, NI), UnitRange(0, NJ))) @@ -868,7 +865,7 @@ def test_as_offset_2d_shift_one_keep_other(): off_arr = np.asarray([[1, 1, 0], [0, 0, 1], [1, -1, 0], [-1, 0, -1]], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == dom i = np.arange(NI)[:, None] @@ -878,7 +875,6 @@ def test_as_offset_2d_shift_one_keep_other(): def test_as_offset_boundary_narrows_domain(): # A uniform out-of-bounds shift narrows the result to the contiguous in-range sub-domain. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -887,7 +883,7 @@ def test_as_offset_boundary_narrows_domain(): np.full(10, -1, dtype=int), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) ) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(1, 10),)) assert np.all(result.ndarray == f.ndarray[0:9]) # out[i] == f[i - 1] @@ -895,7 +891,6 @@ def test_as_offset_boundary_narrows_domain(): def test_as_offset_scattered_oob_raises(): # An out-of-bounds shift in the interior cannot yield a contiguous domain. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -906,12 +901,11 @@ def test_as_offset_scattered_oob_raises(): ) with pytest.raises(ValueError, match="non-contiguous"): - f.premap(as_offset(Ioff, off)) + f.premap(as_offset(I, off)) def test_as_offset_introduces_dimension(): # `off` carries a dim the field lacks: the result gains it, out[i, j] == f[i + off[i, j]]. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -921,7 +915,7 @@ def test_as_offset_introduces_dimension(): off_arr, domain=common.Domain(dims=(I, J), ranges=(UnitRange(0, 10), UnitRange(0, 3))) ) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I, J), ranges=(UnitRange(0, 10), UnitRange(0, 3))) assert np.all(result.ndarray == f.ndarray[np.arange(10)[:, None] + off_arr]) @@ -929,14 +923,13 @@ def test_as_offset_introduces_dimension(): def test_as_offset_nonzero_origin(): # Field and offset over a domain that does not start at 0: indices must be shifted by the domain start. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) dom = common.Domain(dims=(I,), ranges=(UnitRange(2, 12),)) f = common._field(np.arange(10).astype(float), domain=dom) off_arr = np.asarray([1, 0, -1, 0, 1, 0, -1, 0, 1, 0], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == dom assert np.all(result.ndarray == f.ndarray[np.arange(10) + off_arr]) @@ -944,7 +937,6 @@ def test_as_offset_nonzero_origin(): def test_as_offset_2d_shift_second_axis(): # Shift along J (the non-leading axis) by a per-(i, j) offset, leave I: out[i, j] == f[i, j + off[i, j]]. - Joff = fbuiltins.FieldOffset("Joff", source=J, target=(J,)) NI, NJ = 3, 4 dom = common.Domain(dims=(I, J), ranges=(UnitRange(0, NI), UnitRange(0, NJ))) @@ -952,7 +944,7 @@ def test_as_offset_2d_shift_second_axis(): off_arr = np.asarray([[1, 1, 0, -1], [0, 0, 1, -1], [1, -1, 0, 0]], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Joff, off)) + result = f.premap(as_offset(J, off)) assert result.domain == dom i = np.arange(NI)[:, None] @@ -960,25 +952,13 @@ def test_as_offset_2d_shift_second_axis(): assert np.all(result.ndarray == f.ndarray[i, j + off_arr]) -def test_as_offset_non_cartesian_offset_raises(): - # `as_offset` only supports Cartesian (self-shift) offsets: single target equal to source. - - off_I = common._field( - np.zeros(3, dtype=int), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 3),)) - ) +def test_as_offset_local_dimension_raises(): + # `as_offset` shifts along a non-local dimension. off_V = common._field( np.zeros(3, dtype=int), domain=common.Domain(dims=(Vertex,), ranges=(UnitRange(0, 3),)) ) - - # 2-element target (neighbor offset) - V2E = fbuiltins.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) - with pytest.raises(ValueError, match="Cartesian"): - as_offset(V2E, off_V) - - # 1-element target but source != target[0] (cross-dim) - IfromJ = fbuiltins.FieldOffset("IfromJ", source=I, target=(J,)) - with pytest.raises(ValueError, match="Cartesian"): - as_offset(IfromJ, off_I) + with pytest.raises(ValueError, match="non-local dimension"): + as_offset(V2EDim, off_V) @pytest.mark.parametrize( 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 9cdb145fca..5efd6c9595 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 @@ -26,33 +26,25 @@ class Dim(gtx.DimensionIndex): ... class LocalDim(gtx.LocalDimensionIndex): ... -CartesianOffset = gtx.FieldOffset("CartesianOffset", source=Dim, target=(Dim,)) - - class UnstructuredOffset(gtx.NeighborConnectivity[Dim, Dim]): Local: typing.TypeAlias = LocalDim def test_domain_deduction_cartesian(): - assert _deduce_grid_type(None, {CartesianOffset}) == gtx.GridType.CARTESIAN assert _deduce_grid_type(None, {Dim}) == gtx.GridType.CARTESIAN + assert _deduce_grid_type(None, {HDim, VDim}) == gtx.GridType.CARTESIAN def test_domain_deduction_unstructured(): assert _deduce_grid_type(None, {UnstructuredOffset}) == gtx.GridType.UNSTRUCTURED assert _deduce_grid_type(None, {LocalDim}) == gtx.GridType.UNSTRUCTURED - # source and target share `.value` but differ in `.kind` -> not Cartesian - CrossKindOffset = gtx.FieldOffset("CrossKind", source=HDim, target=(VDim,)) - assert _deduce_grid_type(None, {CrossKindOffset}) == gtx.GridType.UNSTRUCTURED - # LOCAL self-loop is unstructured - LocalSelfOffset = gtx.FieldOffset("LocalSelf", source=LocalDim, target=(LocalDim,)) - assert _deduce_grid_type(None, {LocalSelfOffset}) == gtx.GridType.UNSTRUCTURED def test_domain_complies_with_request_cartesian(): - assert _deduce_grid_type(gtx.GridType.CARTESIAN, {CartesianOffset}) == gtx.GridType.CARTESIAN - with pytest.raises(ValueError, match="unstructured.*FieldOffset.*found"): + assert _deduce_grid_type(gtx.GridType.CARTESIAN, {Dim}) == gtx.GridType.CARTESIAN + with pytest.raises(ValueError, match="NeighborConnectivity.*local dimension was found"): _deduce_grid_type(gtx.GridType.CARTESIAN, {UnstructuredOffset}) + with pytest.raises(ValueError, match="NeighborConnectivity.*local dimension was found"): _deduce_grid_type(gtx.GridType.CARTESIAN, {LocalDim}) @@ -62,6 +54,4 @@ def test_domain_complies_with_request_unstructured(): == gtx.GridType.UNSTRUCTURED ) # unstructured is ok, even if we don't have unstructured offsets - assert ( - _deduce_grid_type(gtx.GridType.UNSTRUCTURED, {CartesianOffset}) == gtx.GridType.UNSTRUCTURED - ) + assert _deduce_grid_type(gtx.GridType.UNSTRUCTURED, {Dim}) == gtx.GridType.UNSTRUCTURED 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 65a421c171..18358ac418 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 @@ -59,15 +59,9 @@ class Local(gtx.LocalDimensionIndex): ... class TDim(gtx.DimensionIndex): ... -TOff = gtx.FieldOffset("TDim", source=TDim, target=(TDim,)) - - -#: 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.LocalDimensionIndex): ... - - -renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) +#: A connectivity reached through a name other than its declaration's. Lowering must emit the +#: local dimension's tag, not the variable name. +renamed_v2e = V2E class UDim(gtx.DimensionIndex): ... @@ -167,7 +161,7 @@ def foo_float(inp: gtx.Field[[TDim], float64]): def test_as_offset(): def foo(inp: gtx.Field[[TDim], float64], offset: gtx.Field[[TDim], int]): - return inp(as_offset(TOff, offset)) + return inp(as_offset(TDim, offset)) parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) @@ -803,7 +797,7 @@ def foo(edge_f: gtx.Field[gtx.Dims[Edge], float64]): def test_unstructured_shift_lowering_emits_offset_tag_not_variable_name(): - """The IR shift tag is the offset's tag, not the variable the offset is bound to.""" + """The IR shift tag is the local dimension's tag, not the variable it is reached through.""" def foo(edge_f: gtx.Field[[Edge], float64]): return edge_f(renamed_v2e[1]) @@ -811,7 +805,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) - reference = im.as_fieldop(im.lambda_("__it")(im.deref(im.shift("RenamedTag", 1)("__it"))))( + reference = im.as_fieldop(im.lambda_("__it")(im.deref(im.shift(V2EDim.tag, 1)("__it"))))( "edge_f" ) @@ -825,7 +819,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) - reference = im.as_fieldop_neighbors("RenamedTag", "edge_f") + reference = im.as_fieldop_neighbors(V2EDim.tag, "edge_f") assert lowered.expr == reference 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 c21b88be36..881372c568 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 @@ -21,7 +21,6 @@ LocalDimensionIndex, NeighborConnectivity, Field, - FieldOffset, astype, broadcast, errors, @@ -556,39 +555,33 @@ def return_undefined(): def test_as_offset_dim(): - Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) - def as_offset_dim(a: Field[[ADim, BDim], float], b: Field[[ADim], int]): - return a(as_offset(Boff, b)) + return a(as_offset(BDim, b)) with pytest.raises(errors.DSLError, match=f"not in list of offset field dimensions"): _ = FieldOperatorParser.apply_to_function(as_offset_dim) def test_as_offset_dtype(): - Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) - def as_offset_dtype(a: Field[[ADim, BDim], float], b: Field[[BDim], float]): - return a(as_offset(Boff, b)) + return a(as_offset(BDim, b)) with pytest.raises(errors.DSLError, match=f"expected integer for offset field dtype"): _ = FieldOperatorParser.apply_to_function(as_offset_dtype) -def test_as_offset_non_cartesian(): +def test_as_offset_non_dimension(): def as_offset_neighbor(a: Field[[Edge], float], b: Field[[Edge], int]): return a(as_offset(V2E, b)) - with pytest.raises(errors.DSLError, match="Cartesian"): + with pytest.raises(errors.DSLError, match="Expected 1st argument to be of type"): _ = FieldOperatorParser.apply_to_function(as_offset_neighbor) - IfromJ = FieldOffset("IfromJ", source=IDim, target=(JDim,)) - - def as_offset_cross_dim(a: Field[[IDim], float], b: Field[[IDim], int]): - return a(as_offset(IfromJ, b)) + def as_offset_local_dim(a: Field[[Vertex, V2EDim], float], b: Field[[Vertex, V2EDim], int]): + return a(as_offset(V2EDim, b)) - with pytest.raises(errors.DSLError, match="Cartesian"): - _ = FieldOperatorParser.apply_to_function(as_offset_cross_dim) + with pytest.raises(errors.DSLError, match="non-local dimension"): + _ = FieldOperatorParser.apply_to_function(as_offset_local_dim) vpfloat: TypeAlias = float32 diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py index faea3c7764..1a02371357 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py @@ -208,7 +208,7 @@ def verify_testee(): def test_dace_fastcall_with_connectivity(unstructured_case, monkeypatch): """Test reuse of SDFG arguments between program calls by means of SDFG fastcall API.""" - connectivity_E2V = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + connectivity_E2V = unstructured_case.offset_provider[E2V].asnumpy() @gtx.field_operator def testee(a: cases.VField) -> cases.EField: diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index dd58299c02..2d9c31ca1c 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -375,7 +375,7 @@ def testee(a: cases.VField, b: cases.VField): ), ) - SIMPLE_MESH = cases_utils.simple_mesh(None) + SIMPLE_MESH = cases_utils.ir_level(cases_utils.simple_mesh(None)) offset_provider = SIMPLE_MESH.offset_provider test_case = cases.Case.from_mesh_descriptor(SIMPLE_MESH, backend=backend, allocator=backend) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index 873a0f5bb4..467b0b6de7 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -29,6 +29,7 @@ ) from gt4py.next.type_system import type_specifications as ts +from next_tests.integration_tests import cases_utils from next_tests.integration_tests.cases_utils import ( V2E, Edge, @@ -80,7 +81,7 @@ def _translate_gtir_to_sdfg( @pytest.mark.parametrize("has_unit_stride", [False, True]) @pytest.mark.parametrize("disable_field_origin", [False, True]) def test_find_constant_symbols(has_unit_stride, disable_field_origin): - SKIP_VALUE_MESH = skip_value_mesh(None) + SKIP_VALUE_MESH = cases_utils.ir_level(skip_value_mesh(None)) ir = itir.Program( id="find_constant_symbols_sdfg", diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py index a1fb215882..7a9e2e6114 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py @@ -26,6 +26,7 @@ from gt4py.next.iterator.transforms import pass_manager from gt4py.next.type_system import type_specifications as ts +from next_tests.integration_tests import cases_utils from next_tests.integration_tests.cases_utils import ( E2VDim, C2VDim, @@ -72,8 +73,8 @@ def allow_view_arguments(): IOff = im.cartesian_offset(IDim, IDim) # Cartesian shifts are self-describing (`CartesianOffset`), so no offset provider entry is needed. CARTESIAN_OFFSETS: dict = {} -SIMPLE_MESH: MeshDescriptor = simple_mesh(None) -SKIP_VALUE_MESH: MeshDescriptor = skip_value_mesh(None) +SIMPLE_MESH: MeshDescriptor = cases_utils.ir_level(simple_mesh(None)) +SKIP_VALUE_MESH: MeshDescriptor = cases_utils.ir_level(skip_value_mesh(None)) SIZE_TYPE = ts.ScalarType(ts.ScalarKind.INT32) FSYMBOLS = dict( **{gtx_dace_args.range_start_symbol("w", IDim).name: 0}, diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index a5d33b241c..69072d70b2 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -46,6 +46,10 @@ class Local(LocalDimensionIndex): ... class LsqCoeff(LocalDimensionIndex, size=3): ... +class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local = V2E.Local + + def _declare(source: str) -> dict: """ Run `source` as the body of a throwaway module. @@ -387,10 +391,6 @@ def test_from_value_is_an_offset(self): 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_neighbor_index_accepts_numpy_integers(self): from gt4py.next import constructors, embedded @@ -399,14 +399,9 @@ def test_neighbor_index_accepts_numpy_integers(self): codomain=Edge, data=np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), ) - with embedded.context.update(offset_provider={V2E.Local.tag: table}): + with embedded.context.update(offset_provider={V2E: 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 @@ -454,6 +449,10 @@ def test_grid_type_deduction(self): transform_utils._deduce_grid_type(common.GridType.CARTESIAN, [V2E]) +def _table(domain=(Vertex, V2E.Local), codomain=Edge, data=((0, 1, 2, 3), (1, 2, 3, 0))): + from gt4py.next import constructors + + 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 @@ -496,26 +495,65 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): with pytest.raises(TypeError, match="not a connectivity declaration"): common.local_dimension_of(NeighborConnectivity) + data = np.array(data) + return constructors.as_connectivity( + domain=dict(zip(domain, data.shape)), + codomain=codomain, + data=data, + skip_value=common._DEFAULT_SKIP_VALUE, + ) -class TestFieldOffsetDeprecation: - def test_unstructured_field_offset_warns(self): - from gt4py.next import FieldOffset - - with pytest.warns(DeprecationWarning, match="NeighborConnectivity"): - FieldOffset(V2E.Local.tag, source=Edge, target=(Vertex, V2E.Local)) - - def test_derived_and_cartesian_field_offsets_do_not_warn(self, recwarn): - from gt4py.next import FieldOffset - class_ns = _declare( - """ - class C2E(NeighborConnectivity[Vertex, Edge]): - class Local(LocalDimensionIndex): ... - """ - ) - class_ns["C2E"].__gt_field_offset__() - FieldOffset("Koff", source=KDim, target=(KDim,)) - assert not [w for w in recwarn if issubclass(w.category, DeprecationWarning)] +class TestOffsetProvider: + def test_class_keys_become_tags(self): + table = _table() + assert common.as_tag_keyed_offset_provider({V2E: table}) == {V2E.Local.tag: table} + + def test_tag_keys_pass_through(self): + provider = {V2E.Local.tag: _table()} + assert common.as_tag_keyed_offset_provider(provider) is provider + + def test_bare_name_is_rejected(self): + with pytest.raises(TypeError, match="keyed by 'NeighborConnectivity' declarations"): + common.as_tag_keyed_offset_provider({"V2E": _table()}) + + def test_non_string_key_is_rejected(self): + with pytest.raises(TypeError, match="keyed by 'NeighborConnectivity' declarations"): + common.as_tag_keyed_offset_provider({V2E: _table(), Vertex: _table()}) + + def test_binding_twice_is_rejected(self): + with pytest.raises(ValueError, match="twice"): + common.as_tag_keyed_offset_provider({V2E: _table(), V2E.Local.tag: _table()}) + + def test_check_accepts_matching_tables(self): + common.check_offset_provider({V2E.Local.tag: _table()}) + common.check_offset_provider({V2E: _table().__gt_type__()}) + + def test_check_rejects_mismatching_tables(self): + with pytest.raises(ValueError, match="does not match its declaration"): + common.check_offset_provider({V2E.Local.tag: _table(codomain=Vertex)}) + + def test_sharing_connectivity_is_keyed_and_checked_by_its_own_tag(self): + provider = common.as_tag_keyed_offset_provider({V2E: _table(), V2EShared: _table()}) + assert set(provider) == {V2E.Local.tag, V2EShared.tag} + common.check_offset_provider(provider) + with pytest.raises(ValueError, match="'V2EShared' does not match its declaration"): + common.check_offset_provider({V2EShared.tag: _table(codomain=Vertex)}) + + def test_sharing_connectivities_need_the_same_skip_positions(self): + owner = _table(data=((0, 1, 2, 3), (1, 2, 3, 0))) + consistent = _table(data=((3, 2, 1, 0), (0, 3, 2, 1))) + inconsistent = _table(data=((3, 2, 1, -1), (0, 3, 2, 1))) + common.check_offset_provider({V2E: owner, V2EShared: consistent}) + with pytest.raises(ValueError, match="different neighbor structure"): + common.check_offset_provider({V2E: owner, V2EShared: inconsistent}) + + def test_check_skips_undeclared_tags(self): + common.check_offset_provider({"some.hand.written.tag": _table()}) + + def test_missing_connectivity_error_is_actionable(self): + with pytest.raises(KeyError, match="keyed by 'NeighborConnectivity' declarations"): + common.get_offset({}, V2E.Local.tag) class TestConnectivityKeyOver: diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 45ddb1037f..e3f47b0f1e 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -318,7 +318,7 @@ 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]" + main:17:15: note: "Field[Dims[Vertex, Local], float].__call__" has type "def __call__(self, index_field: Connectivity[Any, Any] | type[NeighborConnectivity[Any, Any]], *args: Connectivity[Any, Any] | type[NeighborConnectivity[Any, Any]]) -> Field[Any, Any]" - case: neighbor_connectivity_generic_local main: | From b9e474c90a7af2f3ecd4a0a486216252d81a4049 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 05:39:21 +0200 Subject: [PATCH 11/17] fix[next]: address review of class-keyed offset providers - resolve memoizes only where the module path ends, so a redefined declaration (a re-run notebook cell) resolves to the new class; mismatch messages say so - normalize the provider in DaCe get_sdfg_conn_args and FieldOperatorFromFoast - compare skip positions of shared local dimensions on the tables' device - migration script: import aliases, unqualified names (imports rewritten), bare LOCAL, shadowed names, __all__, tags differing from the variable, Cartesian provider keys reported as removable - drop the working plan document committed by mistake; clear a stale notebook output showing a removed spelling --- ...ectivities-as-types-implementation-plan.md | 1130 ----------------- .../exercises/6_where_domain_solution.ipynb | 18 +- scripts/python/migrate_connectivities.py | 154 ++- .../python/test_migrate_connectivities.py | 56 + src/gt4py/next/common.py | 64 +- src/gt4py/next/ffront/decorator.py | 4 + src/gt4py/next/otf/compiled_program.py | 3 +- .../runners/dace/sdfg_callable.py | 4 +- .../unit_tests/test_neighbor_connectivity.py | 32 + 9 files changed, 277 insertions(+), 1188 deletions(-) delete mode 100644 docs/development/next/connectivities-as-types-implementation-plan.md diff --git a/docs/development/next/connectivities-as-types-implementation-plan.md b/docs/development/next/connectivities-as-types-implementation-plan.md deleted file mode 100644 index 296424e246..0000000000 --- a/docs/development/next/connectivities-as-types-implementation-plan.md +++ /dev/null @@ -1,1130 +0,0 @@ -# Connectivities as types — implementation plan - -**Status**: **APPROVED** at revision 6 (adversarial review rounds 1–4; round 4 verdict APPROVED) -**Target**: an 8-PR stack on `main`, *alternative to* GridTools/gt4py#2844 -**Proposal**: `egparedes/connectivities-as-types` in GridTools/gt4py_knowledge (PR #32) -**Baseline tree**: `b3c53fa7e` (v1.2.2) - -## 0. Scope and relation to #2844 - -The proposal and #2844 agree on *what a dimension is* (a class, its indices its -instances) and disagree on *what identity a dimension has*. #2844 chose -`(tag, kind)` value equality with an interning registry; the proposal chose -nominal type identity with the tag being the qualified Python name. That single -disagreement propagates into five mechanisms, so the two cannot both land. - -This stack **re-cuts** #2844: it keeps the machinery independent of identity, -drops the machinery that exists only to support value identity, and then builds -the connectivity layer on top. **#2844 is closed, not merged** — that is what -makes this an alternative. Consequences: the new dimension ADR is **0028** (the -ADR directory on `main` ends at **0027**; 0028 exists only inside unmerged -#2844), and there is nothing to supersede. - -### Taken from #2844, unchanged in substance - -| Piece | Where in #2844 | -| -------------------------------------------------------------------------------------------------------------------------------- | ------------------------------------------ | -| `DimensionMeta` metaclass; `I + 1`, `I > 5`, `repr` living on it | `common.py` | -| `DimensionIndex` base: `__slots__ = ("value",)`, `kind` class keyword, `.dim` property | `common.py` | -| `type Dimension = type[DimensionIndex]` as a PEP 695 alias (so `Dimension("I")` raises rather than silently evaluating to `str`) | `common.py` | -| Metaclass `.value` property raising `AttributeError` that points at `.tag` | `common.py` | -| Deletion of `common.NamedIndex` (`.dim` / `.value` move onto the index instance) | `common.py` + ~40 call sites | -| Deletion of the dimension half of `mypy_plugin.py` (`_DimA`..`_AnyDim`); only the mixed-precision hooks remain | `type_system/mypy_plugin.py` | -| The mechanical migration of every declaration, incl. docs, workshop notebooks and `examples/` (which `test_examples` executes) | 131 files: 52 `src/`, 66 `tests/`, 13 docs | -| `xtyping.resolve_annotation` usage at `fbuiltins._type_conversion_helper` (already on `main` via #2841) | — | - -### Dropped from #2844 - -| Piece | Why | -| --------------------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `_DIMENSION_REGISTRY` interning | identity is the type; nothing to intern | -| `copyreg.pickle(DimensionMeta, _reduce_dimension)` — the **blanket** registration on all dimensions | verified: a module-level dimension class pickles by reference with no help (`pickle.loads(pickle.dumps(KDim)) is KDim`). **But a narrow `copyreg` on `StaggeredMeta` is still required** — see §1.5 | -| The `DimensionMeta`-vs-`DimensionMeta` branch of `__eq__` / `__ne__` | becomes `is`. **The `IntegralScalar` overloads (`I == 5` → `Domain`) are kept**, and therefore so is an explicit `__hash__` — see §1.0 | -| `DimensionIndex.__eq__` comparing `type(self) == type(other)` | becomes `type(self) is type(other)` | -| `common.dimension(tag, kind)` factory | replaced by `common.resolve(tag)`, which imports | -| `fingerprinting.py` `DimensionMeta` deconstructor keyed on `(tag, kind)` | under type identity a dimension *is* fingerprinted by qualified name, so the generic `type` deconstructor is correct — **for the lenient variant only**. The STRICT variant rejects `Staggered[KDim]`, which is not importable under its qualified name. Both in-tree fingerprinters are lenient (`ffront/stages.py:62`, `iterator/ir.py:26`), and `eve_utils.content_hash` (`compiled_program.py:420`) is pickle-based and so needs §1.5's `copyreg`. Record the STRICT caveat in the ADR | -| ADR 0028 as drafted in #2844 | never lands; this stack writes its own 0028 | - -### Changed relative to #2844 - -| Piece | #2844 | This stack | -| --------------------------------- | ------------------------------------------------- | --------------------------------------------------------------------------------------------------------- | -| `tag` default | `cls.__name__`, settable in the class body | `f"{cls.__module__}.{cls.__qualname__}"`, a metaclass property; a class-body `tag = ...` is a `TypeError` | -| Rebuilding a dimension from a tag | `dimension(tag, kind)` (registry) | `resolve(tag)` (`import_module` + `qualname` walk), memoized | -| Declaration site requirement | none | module level, or unpicklable; `` heuristic in `__init_subclass__` | -| Backend name mangling | `tag` used directly | `codegen_name(tag)` + inverse, at ~19 enumerated sites in two name spaces (§1.3(b), (c)) | -| Staggered dimensions | `_Staggered` prefix through the interning factory | `Staggered[D]`, a real parametrized type — **required in PR 2**, not optional (§1.5) | - -### Superseded - -- **#2845** (`test[next]: adopt class-style dimension declarations`) is subsumed - by PR 2: because `dimension()` is not user-facing, the minimal - `I = gtx.dimension("I")` form does not exist and every declaration takes class - form immediately. #2845's pyright coverage is folded in. -- The `FieldOffset`-as-frontend-identifier part of **ADR 0019**. -- **ADR 0026**'s `_Staggered` name prefix (PR 2). - -## 1. Design questions closed before implementation - -Everything in this section was verified by running it, not by reading. Probe -files are named; they become committed test material in the PR that needs them. - -### 1.0 Metaclass mechanics that are easy to get wrong - -**`__hash__` must be declared explicitly.** Python sets `__hash__ = None` on any -class body that defines `__eq__` without `__hash__` — metaclasses included. Since -the `I == 5` → `Domain` overload keeps `__eq__` on `DimensionMeta`, dropping -#2844's `__hash__` makes every dimension class *unhashable*: - -``` ->>> class M(type): -... def __eq__(cls, o): return True ->>> M.__hash__ is None -True ->>> class C(metaclass=M): pass ->>> hash(C) -TypeError: unhashable type: 'M' -``` - -That would break `domain({I: 2})` (`common.py:672-690`), -`Counter[common.Dimension]` (`embedded/nd_array_field.py:314`), -`dict[Dimension, SymbolicRange]` (`iterator/ir_utils/domain_utils.py:136,152`), -`seen: dict[Dimension, Dimension]` (`common.py:1351`), and eve's validator -memoization on annotation objects (`eve/type_validation.py:599`) — so -`ts.DimensionType` would fail at *import*. **Fix: `__hash__ = type.__hash__` -explicitly on `DimensionMeta`,** and likewise on `ConnectivityMeta` if it ever -defines `__eq__`. - -**A metaclass `__getitem__` shadows `__class_getitem__`.** `ConnectivityMeta` -needs `__getitem__` for `V2E[1]` (the single-neighbor shift handle that -`FieldOffset.__getitem__` provides today), but metaclass lookup takes precedence -over `Generic.__class_getitem__`, so a naive implementation makes -`NeighborConnectivity[V, E]` in a bases list fail with -`TypeError: tuple expected at most 1 argument, got 3`. - -**Fix, verified clean under `mypy --strict` and pyright 1.1.414 on Python 3.12** -(`/tmp/probe_meta_getitem3.py`): dispatch on the argument type, delegating -non-`int` subscription back to `cls.__class_getitem__`: - -```python -class ConnectivityMeta(type): - __hash__ = type.__hash__ - - @overload - def __getitem__(cls, item: int) -> Connectivity: ... - @overload - def __getitem__(cls, item: Any) -> Any: ... - def __getitem__(cls, item: Any) -> Any: - # `numbers.Integral`, not `int`: `V2E[np.int32(1)]` must not fall through - # to the type-parameter branch (it raises `TypeError: V2E is not a - # generic class` there). `bool` is excluded so `V2E[True]` is an error - # rather than silently neighbor 1. - if isinstance(item, numbers.Integral) and not isinstance(item, bool): - return _bound_single_neighbor(cls, int(item)) - # type-parameter subscription, e.g. `NeighborConnectivity[V, E]` - return cast(Any, cls).__class_getitem__(item) -``` - -`cast(Any, cls)`, not `super()` — `__class_getitem__` is on the class, not on the -metaclass MRO; `super().__class_getitem__` raises `AttributeError`. With the cast -both checkers report zero errors and all uses work at runtime -(`NC[V, E]`, `class V2E(NC[V, E])`, `V2E[1]`, and `V2E.Local` as an annotation). -pyright accepts `NC[V, E]` in a **bases list**; in a *value* position it types it -`Any`, which is why the overloads above matter — without them `V2E[1]` is also -`Any` and the shift handle is untyped. - -### 1.1 `NeighborConnectivity` is **not** a `Connectivity` (proposal Open Q6) - -`common.Connectivity` is `Field[DimsT, IntegralScalar]` — a **data** protocol -(`common.py:990`; `ndarray`, `asnumpy`, `domain` are all on it). A declaration -class holds no data. - -**Resolution.** Two distinct things, distinct hierarchies: - -- `NeighborConnectivity` — a **declaration**. Not a `Connectivity`. It produces a - `NeighborConnectivityType` via `__gt_type__()`, is the provider key, and is the - handle written in DSL code (`a(V2E)`). -- `NeighborTable` / `NdArrayConnectivityField` — the **data**, unchanged, still - `Connectivity` implementations. - -This is the shape `FieldOffset` already has: it is *not* a `Connectivity` either, -and `premap` special-cases it at `nd_array_field.py:317-320`. So `a(V2E)` -continues to work by widening the same union — `Field.premap` and -`Field.__call__` are typed `Connectivity | fbuiltins.FieldOffset` -(`common.py:785, 791-794`) and become `Connectivity | type[NeighborConnectivity]` -in PR 4. `V2E` has **no instances**: `ConnectivityMeta.__call__` raises -`TypeError("… is a connectivity declaration and cannot be instantiated; bind a table through the offset provider")`. - -Consequence: the proposal's sketch line -`class NeighborConnectivity(Connectivity[MultiDimensionIndex[Origin, Local], Codomain], ...)` -is **wrong and dropped**. `MultiDimensionIndex` remains the *domain index type* of -the `NeighborTable` (PR 8). **The knowledge-repo note needs this correction.** - -### 1.2 How `Local` reaches the base (proposal Open Q2) - -`requires-python = '>=3.12'`, so a PEP 696 default type parameter (3.13) is not -available. Resolution: **metaclass discovery**, base carrying a `ClassVar` -annotation, subclass declaring the nested class explicitly: - -```python -class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( - metaclass=ConnectivityMeta -): - Local: ClassVar[type[LocalDimensionIndex]] # annotation only, never assigned - - -class V2E(NeighborConnectivity[V, E], max_neighbors=6): - class Local(LocalDimensionIndex): ... # explicit, required -``` - -Verified under `mypy --strict --python-version 3.12` and `pyright --pythonversion 3.12` -(`/tmp/probe_local.py`): - -| Variant | base declares | mypy | pyright | -| ------- | ------------------------------------------------ | ----- | ---------------------------------------- | -| 1 | `Local: ClassVar[type[LocalDimensionIndex]]` | clean | clean | -| 2 | nothing | clean | clean | -| 3 | a real nested `class Local(LocalDimensionIndex)` | clean | **`reportIncompatibleVariableOverride`** | - -In all three the intended negative case (`Field[V, A.Local]` vs -`Field[V, B.Local]`) is correctly an error. Variant 3 is rejected. - -**Stated precisely — what variant 1 does and does not buy.** It does *not* make -`conn.Local` usable as a **type annotation** when `conn` is a generic -`type[NeighborConnectivity]`: both checkers reject that (mypy `name-defined`, -pyright `reportInvalidTypeForm`), and `T.Local` on a `TypeVar` is rejected too. -What variant 1 buys over variant 2 is only **value-level** access — -`reveal_type(conn.Local)` is `type[LocalDimensionIndex]` instead of an attribute -error — which is what library code in `common`, the backends and -`type_synthesizer` actually needs. Variant 1 is chosen for that, not for generic -annotations. Generic library code that must *name* a local dimension in a -signature uses `type[LocalDimensionIndex]`. - -This extends the proposal's probe P2: a **generated** `Local` is unusable as an -annotation, but a base `ClassVar` *annotation* plus an explicitly declared nested -class is fine. - -### 1.3 The IR keeps string tags; `resolve()` and `codegen_name()` are both required - -`AxisLiteral.value: str` stays (making it carry the class is a separate IR -change, deferred past this stack). It now holds the **qualified** tag, and that -has two consequences the first draft of this plan underestimated. - -**(a) `resolve(tag)` at every rebuild site**, memoized — `inference.py:464` calls -it once per `AxisLiteral` on the type-inference hot path: - -| Site | Purpose | -| ------------------------------------------- | -------------------------------------------------------------------------------------------------------------------------------- | -| `iterator/ir_utils/domain_utils.py` | `AxisLiteral` → `Dimension` | -| `iterator/ir_utils/misc.py` | `AxisLiteral` → `Dimension` | -| `iterator/type_system/inference.py:464` | `AxisLiteral` → `ts.DimensionType` | -| `codegens/gtfn/itir_to_gtfn_ir.py` (×2) | staggered-name sniffing → replaced in PR 2 by `Staggered[D]` | -| `dace/lowering/gtir_to_sdfg_lambda.py:1155` | synthesizes the local dim from the *offset* tag: `Dimension(offset, LOCAL)`. **Must be fixed in PR 2, not deferred** — see below | -| `dace/sdfg_args.py` | axis name → `Dimension` | -| `runners/roundtrip.py` | emits `gtx.Dimension(...)` as *source text* → becomes an import | -| ~~`common.flip_staggered` (×2)~~ | **not** a `resolve()` site: `Staggered[D]` replaces it with an interning subscript, §1.5 | - -`resolve` on a nested qualname was verified to work and round-trip -(`resolve("mymod.V2E.Local") is mymod.V2E.Local`), which matters because PR 4 keys -the provider on `V2E.Local.tag`. **One hazard to settle in PR 2**: a purely dotted -tag does not record *where* the module path ends and the qualname begins, so -`resolve` must try the longest importable prefix and walk the rest — O(depth) -import attempts, and in principle ambiguous if a module path and a class-attribute -chain collide. `pickle` avoids this by storing module and qualname *separately*. -Options: keep the pure dotted form (what the proposal asks for, ambiguity -tolerated and memoized away) or use an explicit separator such as -`"module:qualname"`. **Recommendation: keep the dotted form** — it is what makes -the tag "also a valid tag string for the IR", the collision requires a module and -an attribute chain to have the same spelling, and `resolve` can prefer the -*longest* importable prefix so a real module always wins. Record the residual in -the ADR. - -**(b) `codegen_name(tag)` — dots are illegal in every generated identifier.** -`eve`'s `SymbolName`/`SymbolRef` are constrained by -`_SYMBOL_NAME_RE = ^[a-zA-Z_]\w*$` (`eve/concepts.py:23,26,32`), so a qualified -tag reaching `Sym(id=...)` is a *validation error*, not a cosmetic problem. The -first draft mentioned mangling only in the abstract and put the roundtrip change -in a later PR; both were wrong. All of these are **PR 2**: - -| Site | What breaks without mangling | -| -------------------------------------------------------- | ----------------------------------------------------------------------------------------------------------- | -| `codegens/gtfn/itir_to_gtfn_ir.py:170-195` | `TagDefinition(name=Sym(id=dim.value))` → `SymbolName` validation error | -| `codegens/gtfn/gtfn_module.py:97, 130-136` | `generated::{dim.value}_t`, plus `name.lower()` | -| `otf/binding/nanobind.py:197, 211` | C++ identifiers | -| `dace/lowering/gtir_to_sdfg_utils.py` `get_map_variable` | `i_{dim.value}_gtx_{kind}` → invalid DaCe symbol | -| `dace/sdfg_args.py:80` `_field_symbol` | invalid DaCe symbol | -| `dace/lowering/gtir_python_codegen.py:137-138` | `visit_AxisLiteral` returns the raw value | -| `runners/roundtrip.py:64, 177` | `AxisLiteral = as_fmt("{value}")`, and `{o.value} = gtx.Dimension(...)` emits `a.b.I = ...` → `SyntaxError` | - -**(d) The mangling scheme, corrected.** Earlier drafts said "injective (escape -existing `__` before replacing `.`)", i.e. `_ -> __` then `. -> _`. **That is not -injective**: `.` becomes a single `_`, so `".."` and `"_"` both map to `"__"`. -Exhaustively tested over the alphabet `{a, ., _}` up to length 6 -(`/tmp/probe_mangle.py`): **686 collisions in 1092 inputs.** Since a generated -identifier may only contain `[A-Za-z0-9_]`, `_` is the only available separator -and a *prefix escape* is required: - -```python -def codegen_name(tag: Tag) -> str: - return tag.replace("_", "_u").replace(".", "_d") - - -def from_codegen_name(name: str) -> Tag: - return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) -``` - -Every `_` in the output is the first character of a two-character escape, so -decoding is unambiguous. Verified exhaustively over `{a, ., _, u, d}` up to -length 6 — **19530 inputs, 0 collisions, 0 round-trip failures**, every output a -valid identifier, including the adversarial `"_u"`, `"_d"` and `"a_ud.b"` -(`/tmp/probe_mangle2.py`). Cost: names grow (`mod.V2E.Local` → -`mod_dV2E_dLocal`), which is what gtfn's existing `TagDefinition.alias` mechanism -is for. - -**An inverse is needed too**, wherever generated names are parsed *back* into -dimensions: `dace/sdfg_args.py:25, 60-72` matches `gt_conn_(\S+)` and feeds the -result to `has_offset`. `codegen_name` must therefore be injective *and* have a -`from_codegen_name` partner (escape `__` → `____` before `.` → `__`). - -**A site that cannot be deferred: `gtir_to_sdfg_lambda.py:1155`.** It builds -`gtx_common.Dimension(offset, DimensionKind.LOCAL)` — a local dimension -synthesized from the **offset** tag, which in PR 2 is still a bare provider key -(`"V2E"`) that `resolve()` cannot import. Every DaCe unstructured shift passes -through it, so PR 2 is red on DaCe unless it is fixed there. The fix is local and -available: `conn_type` is already in scope (`:1134-1152`) and `:1135` already -asserts `conn_type.domain[1].kind == LOCAL`, so the line becomes -`offset_type = conn_type.domain[1]` (equivalently `conn_type.neighbor_dim`). -It is *necessary but not sufficient* for PR 1's `shift × tag≠localdim` DaCe cell: -that cell fails earlier, at `gtir_to_sdfg.py:842` -(`neighbor_table_types[dim.value]`, i.e. A4 on the connectivity *argument's* local -dim), before `:1155` is reached — and after the `:1155` fix, `:1371`/`:1455` would -reference `gt_conn_` while `:1104`/`:722` declare `gt_conn_`. So -**the DaCe shift cell stays in the skip matrix until PR 4**, where the -single-string choice makes both agree. (An earlier draft said PR 2; that would -leave PR 2 red on that cell.) - -**(c) The *offset* key is a second dotted name space, and it is mangled in PR 4, -not PR 2.** §1.3(b) covers *dimension* names only. When PR 4 makes the provider -key `cls.tag`, the **offset** string that flows through the IR -(`OffsetLiteral.value`, the provider key, `o` in the gtfn/DaCe connectivity -plumbing) becomes dotted too, and a different set of sites turns *it* into an -identifier. These are all **PR 4**: - -| Site | What breaks | -| ------------------------------------------------------------------------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `codegens/gtfn/itir_to_gtfn_ir.py:184` | `TagDefinition(name=Sym(id=offset_name))` → `SymbolName` regex | -| `codegens/gtfn/itir_to_gtfn_ir.py:490` | `SymRef(id=o)` for each connectivity → `SymbolRef` regex | -| `codegens/gtfn/codegen.py:147-148` | `visit_OffsetLiteral` emits `node.value` raw into C++ | -| `codegens/gtfn/gtfn_module.py:118, 132, 136` | `GENERATED_CONNECTIVITY_PARAM_PREFIX + name.lower()`, `generated::{name}_t` | -| `dace/sdfg_args.py:56` | `connectivity_identifier(name)` → `gt_conn_a.b.V2E`, an invalid SDFG array name | -| `dace/sdfg_args.py:60`, `dace/workflow/bindings.py:200, 286` | `is_connectivity_identifier` / `_parse_gt_connectivities` — the **inverse** direction, so `from_codegen_name` has *several* live consumers, not one | -| `dace/workflow/translation.py:61`, `dace/sdfg_callable.py:103`, `dace/program.py:156` | `connectivity_identifier(offset)` again, on the argument-binding path | -| `dace/lowering/gtir_to_sdfg_lambda.py:1104, 1371, 1455, 1727`, `gtir_to_sdfg.py:722` | the same identifier, consumed in the lowering | -| `runners/roundtrip.py:63` | `OffsetLiteral = as_fmt("{value}")` — emits the offset tag *raw as Python source*, into the program **body**; mangling `:176` alone still leaves `NameError: name 'tests' is not defined` | -| `dace/sdfg_args.py:83-84` | `_field_symbol`: `assert m[1] in offset_provider_type` — a *second* `from_codegen_name` consumer besides `:70` | -| `dace/lowering/gtir_to_sdfg_lambda.py:1892` | `visit_OffsetLiteral` → `SymbolExpr(node.value, INDEX_DTYPE)`, i.e. a dotted string used as a DaCe symbolic expression | -| `runners/roundtrip.py:152, 176` | collects offset-literal strings, then `f'{o} = offset("{o}")'` → `a.b.V2E = offset(...)` → `SyntaxError` | - -So `codegen_name` / `from_codegen_name` are introduced in PR 2 for dimensions and -**applied again in PR 4 for offsets**, at ~16 further sites. Two earlier claims -were wrong: that `from_codegen_name`'s only live consumer is in PR 2, and that the -DaCe surface is confined to `sdfg_args.py` and the lowering — the -argument-binding path (`workflow/translation.py`, `workflow/bindings.py`, -`sdfg_callable.py`, `program.py`) carries it too, in both directions. - -Because of (b) and (c), the review shortcut "diff PR 2 against #2844, the delta is -only identity" is **false**: #2844 needed none of this. Reviewers should expect a real -mangling layer on top of the identity delta. - -**`AxisLiteral.kind` becomes redundant** (the class carries it) — the `TODO` at -`iterator/ir.py:93`. Kept in PR 2, removed in PR 7, to keep PR 2's IR-expectation -churn to the `value` strings only. - -### 1.4 `LocalDimensionIndex` subclasses `DimensionIndex`; `DimensionBaseIndex` is dropped - -The proposal lists `DimensionBaseIndex` as a separate root with `DimensionIndex` -and `LocalDimensionIndex` as siblings. That does not survive contact with the -tree: `Dimension` is `type[DimensionIndex]`, eve validates a `type[X]` -annotation by `issubclass` (verified: a subclass passes, the base and an -unrelated class are both rejected), so sibling local dimensions would force -widening to `type[DimensionBaseIndex]` at `ts.DimensionType.dim`, -`ts.FieldType.dims`, `ConnectivityType.domain`, `Domain.__init__` and the `DimT` -/ `DimT_co` bounds — and would then accept local dimensions everywhere a primary -one is meant, which is the same looseness with extra ceremony. - -**Resolution**: `class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL)`. -`DimensionBaseIndex` is not introduced at all — one concept fewer, which is the -proposal's own stated goal. Where primary-only is required the check is -`dim.kind is not DimensionKind.LOCAL`, exactly as today. This also removes -#2844's deferral note ("a `DimensionBase` root, deferred until the requirements -of non-user-declarable dimensions are known") as a thing that needs resolving. - -**Verified**: all 38 sites in `src/` that discriminate a local dimension do so by -a **runtime `kind` check**, not by a static type distinction -(`transform_utils.py:65`, `type_deduction.py:460, 774`, -`custom_layout_allocators.py:171`, `past_to_itir.py:409`, `common.py:1168, 1336`, -`nd_array_field.py:972, 976`, `gtfn_module.py:91`, `embedded.py:922`, -`gtir_to_sdfg_types.py:76`, …). The tree already treats local dimensions as -`Dimension`s everywhere — `ConnectivityType.domain: tuple[Dimension, ...]` -includes the local one — so subclassing loses nothing it currently relies on, and -`Dims` (`tuple[Unpack[ShapeTs]]`, `common.py:57`) puts no bound on its members -either. - -**What subclassing does cost**, and the mitigation: every `DimensionIndex` -*bound* now statically admits a local dimension — -`NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]`, -`Staggered[D: DimensionIndex]` and `MultiDimensionIndex[D: DimensionIndex, *Ls]` -would all accept `V2E.Local` as their primary parameter. Each therefore gets a -runtime `kind is not DimensionKind.LOCAL` check in `__init_subclass__` / -`__class_getitem__`, and `LocalDimensionIndex.__init_subclass__` rejects an -explicit `kind=` other than `LOCAL`. This is the same runtime-check discipline -the tree already uses; the static gap is the price of the concept removed. - -**Deviation from the proposal; needs feeding back to the note.** - -### 1.5 `Staggered[D]` is required in PR 2, not PR 7 - -`flip_staggered` builds `Dimension(f"_Staggered{name}")` from a string -(`common.py:1452-1457`) and `is_staggered` tests `dim.value.startswith(prefix)` -(`:1447-1449`). #2844 routes both through the interning factory. With the -registry gone there is **no importable `_Staggered` type**, and a -dynamically created class would get the tag -`gt4py.next.common._Staggered`, so `is_staggered` is false and -`as_non_staggered` cannot recover the base dimension's module. Live dependents: -`test_staggered.py` (233 lines), `cases_utils.py:161` -(`KHalfDim = flip_staggered(KDim)`), gtfn `_add_staggered_aliases` -(`itir_to_gtfn_ir.py:203-215`), DaCe `get_map_variable` -(`gtir_to_sdfg_utils.py:52`), `type_synthesizer`, `test_common.py`, -`test_domain_utils.py`. - -So PR 2 is **not green** without `Staggered[D]`. It is Cartesian-only and does -not depend on the connectivity layer, so it moves into PR 2. - -**The obvious mechanism does not work.** A PEP 695 generic -`class Staggered[D: DimensionIndex](DimensionIndex)` makes `Staggered[KDim]` a -`typing._GenericAlias`, **not a class** (verified, `/tmp/probe_staggered.py`): - -``` -type(Staggered[KDim]) -> -isinstance(Staggered[KDim], type)-> False -issubclass(Staggered[KDim], ...) -> TypeError: issubclass() arg 1 must be a class -Staggered[KDim].tag -> '__main__.Staggered' # KDim is gone -``` - -So it fails eve's `type[DimensionIndex]` validation and its tag cannot name the -base dimension — it is not a `Dimension` at all. - -**The mechanism that does work** (verified, `/tmp/probe_staggered3.py`: runs -correctly and is **0 errors under both `mypy --strict` and pyright 1.1.414** on -3.12) is a metaclass `__getitem__` that *builds and interns a real class*, paired -with a `TYPE_CHECKING` declaration so checkers still see an ordinary generic: - -```python -class StaggeredMeta(DimensionMeta): - def __getitem__(cls, base: Dimension) -> Dimension: - if base not in _staggered_cache: - _staggered_cache[base] = StaggeredMeta( - f"Staggered[{base.__name__}]", - (cls,), # NOT (cls, base) -- see below - { - "_tag": f"{cls.__module__}.{cls.__qualname__}[{base.tag}]", - "kind": base.kind, - "base": base, - "__slots__": (), - }, - ) - return _staggered_cache[base] - - -if TYPE_CHECKING: - - class Staggered[D: DimensionIndex](DimensionIndex): - base: ClassVar[Dimension] -else: - - class Staggered(DimensionIndex, metaclass=StaggeredMeta): - __slots__ = () - base: ClassVar[Dimension] -``` - -Verified properties of `Staggered[KDim]`: it *is* a class; -`tag == "gt4py.next.common.Staggered[]"`; `kind` is -inherited from the base; `issubclass(_, DimensionIndex)` and -`issubclass(_, Staggered)` hold; it is instantiable as an index; and -`Staggered[KDim] is Staggered[KDim]`, so identity is stable. `Staggered[KDim]` in -an annotation and inside `Field[Dims[Staggered[KDim]], float]` are both accepted -by both checkers. - -- **Bases are `(cls,)`, not `(cls, base)`.** Inheriting from the base dimension - would make `issubclass(Staggered[KDim], KDim)` true, i.e. `KHalfDim` would be - accepted everywhere `KDim` is required. It is a *different* dimension; only - `kind` is inherited, copied explicitly into the namespace. -- `is_staggered(dim)` becomes **`"base" in dim.__dict__`**, not - `issubclass(dim, Staggered)`, and `as_non_staggered(dim)` becomes `dim.base`. - Two runtime facts force this: `issubclass(Staggered, Staggered)` is true for - the bare base, which has no `base`; and `Staggered[KDim]` is *subclassable* - (`class KHalf2(Staggered[KDim])` yields a second, un-interned staggered-K type - with tag `.KHalf2`). `Staggered.__init_subclass__` therefore rejects - any subclass the metaclass did not create, so the interned form is the only - one. Still structural — no string sniffing. -- **The guards were verified, including the escape routes** - (`/tmp/probe_staggered_guards.py`). All four are blocked: - `class KHalf2(Staggered[KDim])`, `class X(Staggered)`, a direct - `StaggeredMeta("Y", (Staggered,), {})`, and the double subscript - `Staggered[KDim][KDim]`. The `copyreg` fallback round-trips the bare - `Staggered` by reference, the parametrized class with identity preserved, and - instances. Implementation note: gate `__init_subclass__` on a **namespace - marker** the metaclass sets (`"_tag" in cls.__dict__`), not on a module-level - "currently building" flag — the flag works but is not thread-safe, and - compilation runs in worker processes and threads. Three further honest limits: - the guard defends against **accidental** subclassing only — a deliberate - `StaggeredMeta("Forged", (Staggered,), {...marker})` or `types.new_class` can - still forge a same-`tag`, non-identical type (as it can for any class); - `Staggered[Staggered[KDim]]` must be rejected explicitly by testing - `"base" in base.__dict__` in `__getitem__`, or it nests and pickles happily; and - a hand-built `copyreg` payload such as `(_make_staggered, (int,))` should raise - a `TypeError` naming the offending base rather than an `AttributeError`. -- `resolve` gains the `[]` grammar: it parses the brackets and - evaluates `Staggered[resolve(inner)]`, which hits the same intern cache, so a - staggered dimension round-trips through the IR to the *same* class object. -- **A narrow `copyreg` is required after all.** `Staggered[KDim]`'s - `__qualname__` is `Staggered[KDim]`, which `pickle.save_global` cannot look up: - `PicklingError: Can't pickle : attribute lookup Staggered[KDim] on … failed` (verified, `/tmp/probe_staggered_pickle.py`). A - `copyreg.pickle(StaggeredMeta, lambda cls: (_make_staggered, (cls.base,)))` - fixes it *and preserves identity*, because the reconstructor goes back through - the intern cache. **But it must guard the bare base**: `type(Staggered) is StaggeredMeta` too, so a reducer that unconditionally reads `cls.base` fails on - `Staggered` itself with `AttributeError: type object 'Staggered' has no attribute 'base'` (verified — an earlier draft of this section claimed the - registration "captures only parametrized dimensions", which is false). The - reducer therefore falls back to by-reference pickling when - `"base" not in cls.__dict__`. It never captures a plain dimension - (`type(KDim) is DimensionMeta`). This is materially narrower than - #2844's blanket registration on `DimensionMeta` — a parametrized type needs a - reconstructor for the same reason `typing` aliases do — but §0's "`copyreg` - dropped" row is only true of the blanket form, and the ADR must say so. -- **Two honest costs.** (i) `_staggered_cache` is a cache, and the proposal's - headline is that the *name-keyed* registry goes away. The difference is real but must be - stated: it is keyed by a *dimension class*, is internal, and is memoization of - a type constructor (as `typing`'s own subscription cache is), not interning of - user-authored name strings — nothing resolves a user string through it. - (ii) the `TYPE_CHECKING` split means the static and runtime definitions can - drift; a unit test must assert the runtime facts the static form does not - express (real class, `issubclass` against `Staggered` but *not* against the - base, tag shape, interning). -- Supersedes ADR 0026, recorded in the PR-2 ADR. - -## 2. The PR stack - -Branches follow the repo's stacked convention, `connectivities-as-types--`, -each based on its predecessor, all targeting `main`. PR titles are Conventional -Commits (squash-merge lands the title). - -______________________________________________________________________ - -### PR 1 — `fix[next]: lower unstructured shifts with the offset's own tag` - -**Independent of the rest of the stack; lands first, on its own merit.** - -`foast_to_gtir._visit_shift` emits the **Python variable name** as the IR shift -tag (`foast_to_gtir.py:305` `offset_name.id`, `:331` `str(offset_name)`), because -`ts.OffsetType` does not carry the tag. Embedded execution keys on -`FieldOffset.value`. So the same program needs a *different* provider key -depending on the backend — confirmed by running it on v1.2.2: - -``` -MyOff = FieldOffset("TAGNAME", ...) -embedded: {"TAGNAME": conn} OK ; {"MyOff": conn} -> KeyError 'TAGNAME' -roundtrip: {"MyOff": conn} OK ; {"TAGNAME": conn} -> KeyError 'MyOff' -``` - -**Change** - -- `ts.OffsetType` gains **`tag: Optional[Tag] = None`** — *not* a required field. - `type_deduction.py:709` builds `ts.OffsetType(source=conn.codomain, target=(conn.domain_dim,))` from `IDim + 1`, a `CartesianConnectivity` that has - no tag at all; making `tag` required breaks it. -- `FieldOffset.__gt_type__` fills it (`fbuiltins.py:485`). -- `type_deduction.py:464`, which rebuilds an `OffsetType` when `Off[1]` drops the - local dimension, must **propagate** the tag. -- `foast_to_gtir._visit_shift`: the `Subscript` branch and the bare `Name` branch - use `arg.type.tag`, asserting non-`None` (both are unstructured paths, where a - tag always exists). - -**Tests.** `tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py` -today covers exactly `a(Off[1])` on `GTFN_CPU`. Extend to -{shift, `neighbor_sum`} × {embedded, roundtrip, gtfn, dace} × {tag≠varname, -tag≠local-dim-name}. - -**The matrix is not uniform, and a blanket `xfail` will not do.** `xfail_strict = true` -(`pyproject.toml:323`), and measured behaviour on v1.2.2 is: - -| case | embedded | roundtrip | gtfn | dace | -| ---------------------------- | -------- | --------- | -------- | ------------------------------------------------------------------------------------------ | -| shift, tag≠varname | pass | pass | pass | pass | -| shift, tag≠localdim | pass | pass | pass | **fail** `KeyError` (`gtir_to_sdfg_lambda.py:1155` synthesizes the local dim from the tag) | -| `neighbor_sum`, tag≠localdim | **fail** | **pass** | **fail** | **fail** | - -So the first draft's acceptance criterion ("shift cells pass on all four -backends") is unreachable before the backend work, and a strict blanket `xfail` -would XPASS on roundtrip. **Fix**: add a per-backend skip matrix entry in -`tests/next_tests/definitions.py` (a new `USES_*` marker) covering exactly the -failing cells, roundtrip excluded. **They are removed in two steps**: the -`shift × tag≠localdim` DaCe cell and the three `neighbor_sum × tag≠localdim` cells -all in **PR 4**, where the single-string choice makes A3/A4 vacuous — *not* in -PR 5, and *not* the DaCe cell in PR 2 (the `:1155` fix there is necessary but not -sufficient; see §1.3(a)). The gtfn `neighbor_sum` -failure is now confirmed **by running it**; the proposal had it only "by -reading". - -**No CHANGELOG entry.** Verified against the history: `CHANGELOG.md` is touched -*only* by release PRs (`git log -- CHANGELOG.md` is release commits exclusively, -and nothing between `b3c53fa7e` and `upstream/main` touches it). The behaviour -change — which key a compiled backend requires when tag ≠ variable name — belongs -in the PR description, and reaches the changelog when the release PR is cut. Two -earlier drafts of this plan said otherwise, including for PR 6's breaking change. - -**ICON4Py is unaffected by PR 1**: all 16 `FieldOffset` variable names equal -their tags (`model/common/src/icon4py/model/common/dimension.py:33-48`). - -**Acceptance**: `nox -s test_next` green; every cell in the matrix either passes -or is covered by the documented skip matrix. - -______________________________________________________________________ - -### PR 2 — `feat[next]: a concrete Dimension is a class, identified by its qualified name` - -The #2844 core with the identity divergences of §0, **plus** the mangling layer -of §1.3(b) and `Staggered[D]` of §1.5 — both of which #2844 did not need and -without which this PR cannot be green. Large and largely mechanical. - -**`src/gt4py/next/common.py`** - -```python -class DimensionMeta(type): - kind: DimensionKind - __hash__ = type.__hash__ # §1.0 — mandatory, not optional - - @property - def tag(cls) -> Tag: ... # f"{cls.__module__}.{cls.__qualname__}" - - # operators as in #2844; __eq__/__ne__ keep only the IntegralScalar overload - # (I == 5 -> Domain); the dim-vs-dim branch is `is`. - - -class DimensionIndex(metaclass=DimensionMeta): - __slots__ = ("value",) - kind: ClassVar[DimensionKind] = DimensionKind.HORIZONTAL - - def __init_subclass__(cls, /, kind=None, **kw): ... - - -# Staggered: an interning metaclass + TYPE_CHECKING split, NOT a PEP 695 -# generic -- see §1.5, where the generic form is shown to be unworkable. - - -def resolve(tag: Tag) -> Dimension: ... # memoized; [] grammar -def codegen_name(tag: Tag) -> str: ... # "_" -> "_u", "." -> "_d" (§1.3(d)) -def from_codegen_name(name: str) -> Tag: ... # the inverse, `_([ud])` -> `_` / `.` - - -type Dimension = type[DimensionIndex] -``` - -- `tag` is a metaclass **property**, so it cannot drift from the type. This makes - a class-body `tag = "C2E"` a **silent no-op** (verified: `C2EDim.tag` stays - `"__main__.C2EDim"` even with `tag = "C2E"` in the body) — and that pattern is - exactly what ICON4Py and #2845 use to rename. `__init_subclass__` therefore - **raises** on `"tag" in cls.__dict__`, naming the class and pointing at the - rename path. -- `__init_subclass__` also rejects `"" in cls.__qualname__`. Neither - necessary nor sufficient (`type("Dyn", ...)` in a function passes; a `del`'d - class passes) — the authoritative check stays pickle's own `save_global`. -- `resolve` raises a `ValueError` naming the tag and the failing import, per - CODING_GUIDELINES. - -**Removals**: `NamedIndex`; `_DimA`..`_AnyDim` and the dimension half of -`mypy_plugin.py`; `_DIMENSION_REGISTRY`; `copyreg`; the `fingerprinting.py` -deconstructor; `_STAGGERED_PREFIX` and its string sniffing. - -**Migration**. Every `Dimension("X")` becomes `class X(DimensionIndex): ...` at -module level. Verified counts: 333 `Dimension("` declarations in `tests/`, of -which **133 are function-local across 15 files** and must move to module level; -a dimension *named* `"I"` is declared 46 times across **17** files (the first -draft said 45 files — that was the proposal's *`IDim` file* count, a different -number). Docs, workshop notebooks and `examples/` are included because -`test_examples` executes them; notebook *code* cells only, stored outputs -untouched (they hold recorded tracebacks that must keep naming the symbols that -produced them). - -**IR expectation churn**: 36 `AxisLiteral` and 37 `OffsetLiteral` occurrences in -`tests/`, most already computed from `dim.value`. `test_pretty_roundtrip.py` and -the gtfn/DaCe snapshot tests hold the hardcoded names. - -**Do the sweep with a codemod script, not agent fan-out.** A previous attempt at -agent fan-out on a large mechanical rewrite in this repo died mid-file on the -rate limit and left the tree inconsistent; a script did all 57 files uniformly. - -**ADR 0028** (the directory ends at 0027): nominal identity; the module-level -declaration requirement; `Staggered[D]` superseding ADR 0026; that cache -fingerprints now shift when a declaration moves module (a consequence for ADR -0023, not a reversal); that `resolve()` imports modules named in the IR, which is -the same trust level as `pickle` loading a class by reference. - -**Documented limitation**: interactive `__main__` (REPL, notebooks, `python -c`) -cannot be resolved. `spawn` compile workers re-execute the main *script* as -`__mp_main__`, so file-based `__main__` resolves provided the script has the -`if __name__ == "__main__":` guard the pool already requires. - -**ICON4Py migration script is a PR-2 deliverable, not PR 6.** All 15 local -dimensions and `KDim`/`EdgeDim`/`CellDim`/`VertexDim` have variable name ≠ tag -(`EdgeDim = Dimension("Edge")`), so PR 2 changes every generated symbol and every -cache key downstream. - -**Acceptance**: `nox -s test_next` on **3.12, 3.13 and 3.14** (the `typing` -subscription cache behaves differently per interpreter and this change moves -exactly that behaviour), then `test_eve`, `test_storage`, `test_cartesian`, -`test_examples`; `uv run mypy src/`; `uv run pyright`; `uv run tach check`; -`uv run pre-commit run -a`. One at a time, pytest capped at `-n 4`. - -______________________________________________________________________ - -### PR 2 addendum — the single string moves forward from PR 4 (found during implementation) - -**A gap all four review rounds missed.** Under PR 2 a local dimension's tag becomes its -qualified class name, so `V2EDim = Dimension("V2E", kind=LOCAL)` becomes -`class V2EDim(...)` with tag `mod.V2EDim` — it loses the `"V2E"` spelling that made it -equal to the `FieldOffset` tag and the provider key. **24 of the 38 local-dimension -declarations in `tests/` have a variable name different from their string name**, so PR 2 -by itself breaks constraint A3 (reductions key the provider on the local dimension's name, -`nd_array_field.py:983`, `unroll_reduce.py:47`) and A4 (sparse arguments). The review -checked PR 2 against "conforming programs" and never asked how the test tree would conform -once the tags are qualified. - -**Resolution: PR 4's single-string invariant is pulled forward into PR 2**, with the -existing `V2EDim` class playing the role `V2E.Local` plays later. Every declaration becomes - -```python -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... - - -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -offset_provider = {V2EDim.tag: v2e_table} -``` - -so the FieldOffset tag, the local dimension's tag and the provider key are *one* string -and every path agrees: shifts (PR 1 emits `FieldOffset.value`), reductions and sparse -arguments (the local dim's tag), gtfn's `neighbor_dim`, and DaCe's -`connectivity_identifier(offset_type.tag)`. The #1789 branch becomes dead in PR 2. - -**This reduces total churn rather than adding to it**, provided the references are written -*symbolically*. A use site that says `V2EDim.tag` — not the literal `"tests.….V2EDim"` — -needs no edit in PR 4: PR 4 changes only the declaration, `V2EDim = V2E.Local`, and every -`V2EDim.tag` follows. So the 111 ITIR-level string sites and the provider literals are -rewritten **once, in PR 2, to symbolic `.tag`**, and not again in PR 4. - -The codemod therefore has three jobs, not one: declarations to classes; a `FieldOffset`'s -tag to its *local* dimension's `.tag`; and string provider keys / ITIR offset strings to the -matching symbolic `.tag`. Validate on the central fixtures (`toy_connectivity.py`, -`cases_utils.py`) across every backend before running it tree-wide. - -______________________________________________________________________ - -### PR 3 — `feat[next]: NeighborConnectivity declarations and local dimensions that know their owner` - -**Purely additive**: new concepts next to `FieldOffset`, nothing removed, no -behaviour change, no test churn. - -```python -class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): # §1.4 - owner: ClassVar[type[NeighborConnectivity] | None] = None - max_neighbors: ClassVar[int | None] = None - min_neighbors: ClassVar[int | None] = None - - def __init_subclass__(cls, *, size: int | None = None, **kw): ... - - -class ConnectivityMeta(type): # §1.0 for __hash__ and __getitem__ - @property - def tag(cls) -> Tag: ... - def __call__(cls, *a, **kw) -> NoReturn: ... - - -class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( - metaclass=ConnectivityMeta -): - Local: ClassVar[type[LocalDimensionIndex]] - - def __init_subclass__(cls, *, max_neighbors=None, min_neighbors=None, **kw): ... -``` - -- `__init_subclass__` asserts `"Local" in cls.__dict__` and that it subclasses - `LocalDimensionIndex`, then sets `Local.owner = cls` and copies the counts. A - missing `Local` is a `TypeError` at class creation naming the class. -- **Owner-less** locals: `class LsqUnk(LocalDimensionIndex, size=3)` — `owner is None`, `min == max == size`, never in the provider. ICON4Py's `LsqUnkDim` and - `RBFDimension` need this: they index no table but need sparse storage and - layout. -- Counts are **optional class keywords**, not type parameters (Python has no - integer type parameters and nothing static needs the count). Declared ⇒ a - constraint the table must satisfy. Undeclared ⇒ completed at bind time, from - the table in the JIT flow or from the `NeighborConnectivityType` already passed - through `connectivities=` (`ffront/decorator.py:188-208`) in the AOT flow. - Not static-only because `fvm_nabla_setup.py:99` sizes `V2E` from the atlas - mesh, and ICON skip-value presence is configuration-dependent (`icon.py:130` — - pentagons have skip values on the icosahedron, not the torus). -- **Bind-time validation**, one function replacing constraints A6–A8: shape - `(n, max_neighbors)`, integral dtype, skip values present iff - `min_neighbors < max_neighbors`, `domain[0] is Origin`, `codomain is Codomain`. - -**No provider bridge.** The first draft proposed a dual-keyed -(`Tag | type[NeighborConnectivity]`) provider here. Dropped: the provider is -accessed **directly, not through `get_offset`, at 19 sites in 12 `src/` files** -(`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, -`compiled_program.py`, `pass_manager.py`, …) despite the note at -`common.py:1174`, so a bridge would be both invasive and — since nothing would -exercise class keys — untested. Class-keyed providers land in one place, PR 4. - -**Typing tests**: `typing_probe.py` / `probe_local.py` / `probe_meta_getitem3.py` -become real coverage — `typing_tests/test_next.yaml` cases for -`Field[Dims[V, V2E.Local], float]`, a `TypeVar` bound to `LocalDimensionIndex`, -the negative cross-connectivity case, and the `NC[V, E]`-in-bases case of §1.0; -plus the pyright variants. - -**Acceptance**: full suite green with no behaviour change; new unit tests for -declaration errors, owner wiring, owner-less locals and bind-time validation. - -______________________________________________________________________ - -### PR 4 — `feat[next]: declare connectivities as classes; FieldOffset derived from them` - -**The ordering fix.** Revision 2 put class-keyed providers here and the -declaration migration in PR 6. That cannot be green: `Local.owner` only exists if -the user declared a `NeighborConnectivity` class, and a `FieldOffset` written the -old way (`FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim))`) has no class -to point at — so neither the backend work nor a class-keyed provider has anything -to resolve. The declaration migration must come **first**, and the provider key -must stay a string until the backends are through. - -- The **unstructured** `FieldOffset` is **derived**, not authored: - `FieldOffset.from_connectivity(V2E)` (or `V2E.__gt_offset__()`), which fills - `source = Codomain`, `target = (Origin, V2E.Local)` and — critically — - **`value = V2E.Local.tag`, the *local dimension's* tag, not `V2E.tag`.** -- **Why the local dimension's tag and not the connectivity's.** An earlier draft - used `V2E.tag` and claimed PR 4 was green. It is not: A3 and A4 key the provider - on the **local dimension's** name, at `nd_array_field.py:983` - (`get_offset(provider, axis.value)`, whose in-tree comment is literally - `# assumes offset and local dimension have same name`), `unroll_reduce.py:47` - (`arg.type.offset_type.value`), `gtfn_module.py:95`, `gtir_to_sdfg.py:581, 842`, - `iterator/embedded.py:954, 1519`, and - `gtir_to_sdfg_lambda.py:1371, 1455` (`connectivity_identifier(offset_type.value)`). - Today the `V2EDim = Dimension("V2E")` convention makes that string equal to the - offset tag; PR 4 deletes the convention tree-wide, while the `owner` lookup that - replaces it is PR 5. With `value = V2E.tag` every reduction and every sparse-field - argument would break on embedded, gtfn and DaCe simultaneously — the round-1 - matrix row (`neighbor_sum`, tag≠localdim: embedded/gtfn/DaCe fail) would become - the tree's universal state. - Choosing `V2E.Local.tag` instead makes **all four** of A1, A3, A4 and A5 vacuous - at once, because there is then exactly *one* string and the class produces it. - It also makes the #1789 branch at `itir_to_gtfn_ir.py:181-190` - (`if offset_name != connectivity_type.neighbor_dim.value`) dead already in PR 4. - This is preferable to the alternatives — fusing PR 5 into PR 4, or a transient - `owner`-based fallback inside `get_offset` — because it needs no scaffolding: - the string is simply picked correctly, and PR 5 then removes the dependence on a - string at all. -- **The Cartesian `FieldOffset` constructor stays in PR 4.** An earlier draft said - "every declaration becomes a class", which is wrong: `Ioff`, `Koff` and - `EdgeOffset` (`cases_utils.py:163-169`, e.g. - `Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,))`) and ICON4Py's - `Koff`/`KHalfOff` (`dimension.py:47-48`) are single-target and have no - `NeighborConnectivity` to derive from, and their only remaining consumer — - `as_offset` — does not change until PR 6. Restricting the public constructor to - the single-target form keeps PR 4 green; it disappears with `as_offset` in - PR 6. -- **Providers stay keyed on `Tag`**, now `V2E.Local.tag`. Nothing about the key - *mechanism* changes yet, so the 19 direct-access sites are untouched. A1, A3, A4 - and A5 are all dead at this point — there is one string, and the class produces - it. -- `ts.OffsetType` → `ConnectivityType`, produced by the class - (`type_specifications.py:74` TODO). `type_info.py:637, 858` gate `a(V2E)` - deduction on `ts.OffsetType` and follow. `Field.premap` / `Field.__call__` - unions widen to `Connectivity | type[NeighborConnectivity]` (`common.py:785, 791-794`, `1313-1316`). -- **Test-tree migration lands here**: `toy_connectivity.py`, `cases_utils.py`, - `fvm_nabla_setup.py` are the fixture modules everything imports; 37 - `FieldOffset` sites, 42 `DimensionKind.LOCAL` sites. String provider keys keep - working because they are `cls.tag` — but the tags are now *qualified*, so the - 111 ITIR-level string-offset occurrences in 15 files (`im.shift("V2E")`, - `neighbors("…")`, `OffsetLiteral(value="…")`, string-keyed providers) are - rewritten to `V2E.Local.tag` here rather than in PR 6. -- **ICON4Py**: this is the release-visible declaration change. The migration - script written in PR 2 is extended. - -______________________________________________________________________ - -### PR 5 — `refactor[next]: backends resolve connectivities through the local dimension's owner` - -Where A3, A4 and A5 dissolve and PR 1's skip-matrix entries are removed. Green -while providers are still string-keyed, because a backend goes -`local_dim.owner` → `owner.tag` → the existing lookup: the *identity* question is -answered by the owner pointer, and the key is still a string. - -**The owner-less case must be handled, not assumed away.** At PR 5 `_CONST_DIM` -is still a plain `DimensionIndex(kind=LOCAL)` (it becomes `ConstList` only in -PR 7), and `LsqUnk`-style local axes have `owner is None` by design. Every -converted site reads `getattr(dim, "owner", None)` and falls through when it is -`None`. Two sites already guard by accident — DaCe compares against `_CONST_DIM` -first (`gtir_to_sdfg_lambda.py:1314`) and `unroll_reduce` filters -`offset_type is None` — but `gtfn_module.py:91-98` and `nd_array_field.py:981` -have **no** guard, and ICON4Py never exercises the case -(`test_icon.py:220`), so the gap would not show up downstream. - -`unroll_reduce.py:47` (reads `arg.type.offset_type`, which is the local -`Dimension` — now a `LocalDimensionIndex` carrying `owner`, which is exactly the -back-pointer it lacked), `gtfn_module.py:95, 118, 132`, `itir_to_gtfn_ir.py` -(including the `#1789` `offset_name != neighbor_dim.value` branch at `:181-190`, -which becomes dead and goes), `gtir_to_sdfg.py:581`, -`gtir_to_sdfg_lambda.py:766-770, 1155` (the `Dimension(offset, LOCAL)` synthesis -goes), `nd_array_field.py:981-985` (and its -`# assumes offset and local dimension have same name` comment), -`iterator/embedded.py`, and `runners/roundtrip.py`. - -______________________________________________________________________ - -### PR 6 — `feat[next]!: class-keyed offset providers; remove FieldOffset and the string offset API` - -The only breaking PR, and now the only one that touches the provider key. - -- `OffsetProvider*` become - `Mapping[type[NeighborConnectivity], NeighborTable]`; `get_offset` keys on the - class. **Note the `.owner` hop**: because PR 4 made the IR offset tag the *local - dimension's* tag, `resolve(OffsetLiteral.value)` yields a `LocalDimensionIndex`, - not the connectivity — so the class-key lookup is - `resolve(tag).owner`. (The alternative is to switch the IR tag to `V2E.tag` in - this PR; the `.owner` hop is cheaper and keeps the IR stable.) The **19 direct-access sites in 12 `src/` files** - (`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, - `compiled_program.py`, `pass_manager.py`, …) are converted here, and - `common.py:1174`'s "all accesses should go through `get_offset`" either becomes - true or the note goes. `hash_offset_provider_items_by_id` and the - `fingerprinting` dict handling already tolerate class keys once §1.0's - `__hash__` is in place. -- **Removals**: `FieldOffset` entirely (both forms); `runtime.Offset` as its base - (`fbuiltins.py:467-470` TODO); `iterator/runtime.offset("...")` (12 sites in 6 - files, plus `tracing.py:161-162`); the `V2EDim`-next-to-`V2E` convention; - `embedded/context.py` string plumbing; the `gt4py.next.__init__` exports at - `:47, 140`. -- **`as_offset` changes in the same PR.** It is why the Cartesian `FieldOffset` - form cannot go alone: `ffront/experimental.py:17` + - `type_deduction.py:956-967` require one. New signature - `as_offset(KDim, field)`. Used in 5 test modules, the `Ioff`/`Koff`/`EdgeOffset` - fixtures at `cases_utils.py:163-169`, and **40 non-test call sites in - ICON4Py**. -- `transform_utils.py:50-77` and `past_to_itir.py:77` deduce grid type from the - provider and follow. -- **Accepted double churn**: the ~26 `offset_provider={...}` literals are rewritten - twice — `{V2E.Local.tag: t}` in PR 4, `{V2E: t}` here. The alternative is fusing PR 4 - and PR 6, which loses the green boundary. The ITIR string sites do *not* churn - twice: `im.shift(V2E.Local.tag)` written in PR 4 stays correct. - -**ADR 0029**: the connectivities-as-types record — `FieldOffset` removed, the -class-keyed provider, superseding the `FieldOffset` part of ADR 0019. - -**Breaking-change communication**: the PR title carries the Conventional Commits -`!` marker and the ADR records the removal; the changelog entry is written by the -release PR, not here (see PR 1). No deprecation window — an explicit decision: -ICON4Py's provider keys are bare names, so no import-based shim could have -resolved them. - -**Acceptance**: full suite; `test_fvm_nabla` and `test_icon_like_scan` are the -integration canaries. A before/after run of one gtfn and one DaCe program -checking generated-code equivalence modulo names. - -### PR 7 — `refactor[next]: ConstList replaces _CONST_DIM; AxisLiteral drops kind` - -- `_CONST_DIM` (`iterator/embedded.py:220` and - `dace/lowering/gtir_to_sdfg_lambda.py` — two separate declarations, each - internally consistent; 14 references in total) becomes the owner-less - `ConstList(LocalDimensionIndex, size=1)`, generalizing the magic name from size - 1 to size *n*. -- `AxisLiteral.kind` removed (`iterator/ir.py:93` TODO), now that every tag - resolves to a class carrying its kind. - -______________________________________________________________________ - -### PR 8 — `refactor[next]: MultiDimensionIndex and typed embedded positions` - -- `MultiDimensionIndex[D: DimensionIndex, *Ls]` as the index type of a sparse - position and the domain index of a `NeighborTable`. `*Ls` is unconstrained - because `TypeVarTuple` cannot carry a bound; `__init_subclass__` checks at - runtime what the checker cannot. -- `iterator/embedded.py` positions keyed by dimension types instead of name - strings (`embedded.py:574-576`, `597-616`, `941-950`); `SparseTag` removed. - Constraint A9 dissolves. -- Nothing else depends on this; it is last for that reason. - -______________________________________________________________________ - -## 3. Constraint ledger - -| # | Constraint | Retired by | -| -------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | ------------------------------------------------------- | -| A1 | `FieldOffset.value` == provider key | PR 4 (one declaration produces both) | -| — | *all four string-equality constraints below become vacuous in PR 4*, because the class emits a single string (`V2E.Local.tag`); PR 5 removes the dependence on a string at all | PR 4 / PR 5 | -| A2 | Python variable name == provider key | **PR 1** | -| A3 | local dim name == provider key (reductions) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | -| A4 | local dim name == provider key (sparse args) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | -| A5 | `FieldOffset.value` == local dim name | PR 4 (one declaration) | -| A6 | `target[-1]` == connectivity `neighbor_dim` | PR 3 (bind-time check) | -| A7 | `FieldOffset.source` == `codomain` | PR 3 (bind-time check) | -| A8 | `target[0]` == `domain[0]` | PR 3 (bind-time check) | -| A9 | dim name is the iterator-position dict key | PR 8 | -| A10 | dim name round-trips through `AxisLiteral` | structural; PR 2 qualifies it, PR 7 drops `kind` | -| F5/F8/F9 | codegen name formats | PR 2 (`codegen_name` + inverse) | -| S6 | `as_offset` needs a Cartesian `FieldOffset` | PR 6 | -| — | string-keyed provider | PR 6 | - -## 4. Risks - -1. **PR 2 size, and it is not purely mechanical.** ~150 files of migration *plus* - a name-mangling layer and `Staggered[D]`. It cannot be reviewed as "the - #2844 diff plus identity". Mitigation: land the mangling layer and - `Staggered[D]` as reviewable commits *within* the PR, ordered before the - sweep, so the mechanical part is a separate commit. -2. **Function-local dummy dimensions.** 133 declarations in 15 test files must - move to module level. Under nominal identity, two same-named locals that were - silently the same dimension become distinct — each resulting failure is a - real finding, not churn. -3. **`typing` subscription caching.** Under `(name, kind)` equality, - `Field[Dims[I]] is Field[Dims[I2]]` aliases for two distinct same-named - classes — a known residual of the #2844 design, and an argument *for* this - stack. The cache behaves differently per interpreter: verify on 3.12, 3.13 - and 3.14 **via nox**, not `uv run pytest`. -4. **Fingerprint/cache invalidation.** Moving a declaration between modules now - invalidates compiled artifacts. Intended; CHANGELOG + ADR line. -5. **Naming not yet converged** with `havogt/dependent-local-dimensions` - (`Origin`/`Codomain` vs `source_dim`/`neighbor_dim`; `Local` vs `Dim`; - `min_neighbors` vs `has_skip_values`). PR 3 fixes public names. Converge - before PR 3 is *opened*. -6. **`V2E.Local` vs `Local[V2E]`.** The chain proposals' encodings subscript - `Local`, and a `TypeVar` cannot be subscripted for a nested attribute. Their - semantics are unaffected; their static encoding needs rewriting to `C.Local` - plus a protocol for the generic hop-stack case. Knowledge-repo concern, not a - gt4py blocker. -7. **CSCS GPU CI is flaky and opaque.** All jobs failing at the same second means - infrastructure; `cscs-ci run default` as a PR comment reruns it. #2844's CI is - green except that job. -8. **Two mangling passes, two PRs.** `codegen_name` is applied to dimension names - in PR 2 and to offset names in PR 4, at ~19 sites total, several of which - (`Sym`/`SymRef` construction, DaCe array names) fail *loudly* and several of - which (C++ emission, `name.lower()`) fail only in the generated artifact. - Both PRs need a test that a qualified tag survives a real gtfn and a real - DaCe compile, not just lowering. -9. **`resolve()` on the inference hot path** (`inference.py:464`, once per - `AxisLiteral`). Must be memoized from the start, and the memo must be keyed - so a reloaded module does not return a stale class. - -## 4b. Work the earlier drafts did not mention - -- **Public exports.** `gt4py.next.__init__` must export `NeighborConnectivity`, - `LocalDimensionIndex`, `Staggered` and `resolve` (PR 2 for the dimension half, - PR 3 for the connectivity half), and drop `FieldOffset` / `offset` at `:47, 140` - in PR 6. -- **`type_translation.from_value(V2E)`** works only because the - `hasattr(value, "__gt_type__")` branch at `type_translation.py:328` is tested - *before* the `DimensionMeta` branch. That ordering is load-bearing under this - design and currently untested — PR 3 adds a unit test pinning it. -- **`pyright` is not yet a dependency.** `uv run pyright` appears throughout §5 - but pyright is absent from `pyproject.toml` on `main`; #2845 is what adds it. - PR 2 must explicitly fold in #2845's `typing_exports` / pyright dependency-group - change, or §5's pyright step is not runnable. -- **`test_examples` belongs to PR 4 too.** §5 lists it for PR 2 and PR 6; the - docs and notebooks use `FieldOffset`, so PR 4's declaration migration touches - them and must run it. - -## 5. Verification - -Per PR, in this order, **one at a time** on the shared machine, pytest capped at -`-n 4`: - -``` -uv run pre-commit run -a # ruff, mypy, tach, license headers -uv run pyright # static checks the mypy plugin no longer fakes -uv run nox -s "test_next-3.12(...)" # then 3.13, 3.14 for PR 2 -uv run nox -s test_eve test_storage test_cartesian test_examples # PR 2, PR 6 -``` - -Test-first where behaviour changes, per AGENTS.md: PR 1's regression matrix, PR -3's declaration-error and bind-validation units, and PR 4's provider-key tests -are written before the implementation they cover. - -## 6. Feedback owed to the knowledge-repo note - -- `NeighborConnectivity` is not a `Connectivity`; the sketch's base line is wrong - (§1.1) — this closes Open Q6. -- `DimensionBaseIndex` should be dropped; `LocalDimensionIndex` subclasses - `DimensionIndex` (§1.4). -- Open Q2 is closed: metaclass discovery, base `ClassVar` *annotation*, explicit - nested class — with the precise limit of what that buys statically (§1.2). -- `Staggered[D]` is not a late step; it is a precondition for removing the - name-keyed registry (§1.5) — and it cannot be a PEP 695 generic, needs an - identity-keyed intern cache, and needs a narrow `copyreg`. The note's claim - that `copyreg` disappears entirely is therefore too strong. -- The note's "five name spaces" analysis should record that under qualified tags - there are **two** dotted name spaces reaching codegen — dimension tags and - offset tags — each needing its own mangling pass (§1.3(b) and (c)). -- The note's §Staging step 2 ("`NeighborConnectivity` … object-keyed provider; - `FieldOffset` and string keys removed outright") bundles three changes that - must be separated to stay green: the declaration migration has to precede the - backend work (because `Local.owner` only exists once classes are declared), and - the provider *key* has to stay a string until the backends resolve through the - owner. See PR 4/5/6. -- The staging in the note's §Staging (steps 0–8) is superseded by §2 here; in - particular step 4 ("backends, one file at a time") cannot follow step 2, since - the provider key change and the backend lookups are separable but the mangling - layer is needed at the *dimension* step. - -## 7. Open, non-blocking - -- Naming convergence (risk 5). -- Whether a same-`__name__` collision warning is useful or noise (note Q5). -- Whether interactive `__main__` should be detected with a fallback to in-process - compilation and a warning, or merely documented (note Q3). Plan assumes - documented. -- `AxisLiteral.dim` instead of `AxisLiteral.value` — a follow-up after PR 8. - -______________________________________________________________________ - -## 8. As implemented (PRs 3–6), and where it deviates - -Recorded while implementing; the sections above are the plan as reviewed. - -- **PR 3 was made usable in the DSL**, not just declarative: `ConnectivityMeta.__gt_type__` - gives the offset type, `V2E[i]` and `a(V2E)` work in embedded and on all backends, and - `V2E.Local` in DSL code types as the local dimension (one `visit_Attribute` case). No backend - change was needed, because PR 2 had already made the offset tag, the local dimension's tag and - the provider key one string. Also in PR 3: fingerprinting of a declaration by its dimensions and - counts, `FieldOffset.Local`, and `check_neighbor_table` (explicit until PR 6). -- **Shared local dimensions (not in the plan).** The review of PR 4 found that ICON4Py's flattened - sparse offsets (`C2CE`, `E2ECV`, `E2EC`, `C2CEC`) use another connectivity's local dimension, so - the one-owner rule would have made them unmigratable once `FieldOffset` is gone. A declaration - can adopt an owned local dimension (`Local = C2E.Local`); the owner stays `C2E`, and the sharer - is named in the IR by its own tag: `ConnectivityMeta.offset_tag` is `Local.tag` for the owner - and `cls.tag` for a sharer. This reintroduces an offset tag that differs from its local - dimension's tag — exactly PR 1's `tag != local dim` cells. -- **PR 5 was re-scoped** from "backends resolve through `owner`" to "backends support - connectivities sharing a local dimension". With the single string of PR 2 and class-keyed - providers normalized to tags (PR 6), resolving through `owner` would change nothing; what the - backends did lack was a way to find the table over a local dimension that is not keyed by its - tag. `common.connectivity_key_over(provider, local_dim)` does that (the local dimension's tag if - bound, else any table whose neighbor dimension it is), and replaces the local-tag lookups in - embedded reductions, `unroll_reduce`, gtfn sparse arguments, iterator-embedded sparse lists and - positions, and DaCe (`make_field`, reductions, `map_list`, local-dimension array sizes). All of - PR 1's `uses_offset_tag_differing_from_local_dim*` markers are gone in PR 5. -- **PR 6 keeps the internal provider tag-keyed.** Users key providers by the class; every entry - point (`Program.__call__`, `FieldOperator.__call__`, `compile`, `CompilationOptions`, - `embedded.context.update`, iterator `fendef`) normalizes with `as_tag_keyed_offset_provider` - (class → `offset_tag`; a bare, undotted string key is rejected as the removed `FieldOffset` - spelling). Tables are checked against their declarations once per compiled variant - (`check_offset_provider` in `CompiledProgramsPool._compile_variant`) and on embedded calls, not - on every call. The 19 direct-access sites of §PR 6 therefore needed no change, and the IR keeps - naming connectivities by string, consistently with §1.3. -- **`as_offset(KDim, field)`** takes a dimension in PR 6; Cartesian `FieldOffset`s had no other - use (Cartesian shifts are `Dim + i` since before this stack). -- **Not done**: `ts.OffsetType` is not renamed to `ConnectivityType` (it now only types - connectivity declarations and `as_offset`); `iterator.runtime.offset("...")` stays as the - iterator-level API, since the IR names offsets by string. diff --git a/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb b/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb index 5514c1b4f7..194d72bcc1 100644 --- a/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb +++ b/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb @@ -326,23 +326,7 @@ "execution_count": 12, "id": "ed393959", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Test successful\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/var/folders/2b/_2y31vzs4sl_7rngh2yghbpw0000gn/T/ipykernel_56692/57164705.py:4: UserWarning: Field View Program 'program_domain_where': Using Python execution, consider selecting a perfomance backend.\n", - " program_domain_where(a, b, offset_provider={\"Koff\": K})\n" - ] - } - ], + "outputs": [], "source": [ "test_domain_where()\n", "print(\"Test successful\")" diff --git a/scripts/python/migrate_connectivities.py b/scripts/python/migrate_connectivities.py index d8862873b5..304f7401b2 100644 --- a/scripts/python/migrate_connectivities.py +++ b/scripts/python/migrate_connectivities.py @@ -67,6 +67,8 @@ class Module: tree: ast.Module edits: list[Edit] = dataclasses.field(default_factory=list) notes: list[str] = dataclasses.field(default_factory=list) + #: Names of `gt4py.next` the migrated declarations use unqualified, to be imported. + needed: set[str] = dataclasses.field(default_factory=set) @property def lines(self) -> list[str]: @@ -77,7 +79,7 @@ def segment(self, node: ast.AST) -> str: assert segment is not None return segment - def note(self, node: ast.AST, message: str) -> None: + def note(self, node: ast.stmt | ast.expr, message: str) -> None: self.notes.append(f"{self.path}:{node.lineno}: {message}") @@ -98,8 +100,23 @@ def _keyword(call: ast.Call, name: str, position: int) -> ast.expr | None: return call.args[position] if len(call.args) > position else None +_DECLARATION_CALLEES = ("Dimension", "FieldOffset") + + +def _imported_aliases(module: Module) -> dict[str, str]: + """Local names of imported `Dimension` / `FieldOffset`, e.g. `{"FO": "FieldOffset"}`.""" + return { + alias.asname or alias.name: alias.name + for statement in ast.walk(module.tree) + if isinstance(statement, ast.ImportFrom) + for alias in statement.names + if alias.name in _DECLARATION_CALLEES + } + + def _declarations(module: Module) -> Iterator[tuple[ast.Assign, str, ast.Call, str, str]]: """Module-level `name = [prefix.]Dimension(...)` / `FieldOffset(...)` statements.""" + aliases = _imported_aliases(module) for statement in module.tree.body: if ( isinstance(statement, ast.Assign) @@ -107,9 +124,11 @@ def _declarations(module: Module) -> Iterator[tuple[ast.Assign, str, ast.Call, s and isinstance(statement.targets[0], ast.Name) and isinstance(statement.value, ast.Call) and (callee := _callee_name(statement.value)) is not None - and callee[1] in ("Dimension", "FieldOffset") ): - yield statement, statement.targets[0].id, statement.value, *callee + prefix, name = callee + name = aliases.get(name, name) if prefix == "" else name + if name in _DECLARATION_CALLEES: + yield statement, statement.targets[0].id, statement.value, prefix, name def _replace(module: Module, statement: ast.stmt, text: str) -> None: @@ -122,12 +141,17 @@ def _migrate_declarations(module: Module, cartesian: dict[str, str]) -> None: if kind_of_call == "Dimension": kind = _keyword(call, "kind", 1) kind_src = module.segment(kind) if kind is not None else None - if kind_src is not None and kind_src.endswith(".LOCAL"): - text = f"class {name}({prefix}LocalDimensionIndex): ...\n" + if kind_src is not None and kind_src.split(".")[-1] == "LOCAL": + base = "LocalDimensionIndex" + text = f"class {name}({prefix}{base}): ...\n" elif kind_src is None: - text = f"class {name}({prefix}DimensionIndex): ...\n" + base = "DimensionIndex" + text = f"class {name}({prefix}{base}): ...\n" else: - text = f"class {name}({prefix}DimensionIndex, kind={kind_src}): ...\n" + base = "DimensionIndex" + text = f"class {name}({prefix}{base}, kind={kind_src}): ...\n" + if not prefix: + module.needed.add(base) _replace(module, statement, text) continue @@ -141,6 +165,8 @@ def _migrate_declarations(module: Module, cartesian: dict[str, str]) -> None: f"class {name}({prefix}NeighborConnectivity[{origin}, {module.segment(source)}]):\n" f" Local = {local}\n" ) + if not prefix: + module.needed.add("NeighborConnectivity") _replace(module, statement, text) elif len(target.elts) == 1 and module.segment(target.elts[0]) == module.segment(source): # A Cartesian offset has no declaration any more: `Off[i]` is `Dim + i`. @@ -150,6 +176,48 @@ def _migrate_declarations(module: Module, cartesian: dict[str, str]) -> None: module.note(statement, f"'{name}': a cross-dimension offset has no class equivalent.") +def _migrate_imports(module: Module) -> None: + """Drop imports of the removed `FieldOffset`; import the class names used unqualified.""" + fieldoffset_names = { + local for local, name in _imported_aliases(module).items() if name == "FieldOffset" + } + imported = { + alias.asname or alias.name + for statement in ast.walk(module.tree) + if isinstance(statement, ast.ImportFrom) + for alias in statement.names + } + missing = sorted(module.needed - imported) + added = False + for statement in module.tree.body: + if not isinstance(statement, ast.ImportFrom): + continue + names = {alias.asname or alias.name for alias in statement.names} + drops = names & fieldoffset_names + adds_here = bool(missing) and not added and bool(names & {"Dimension", "DimensionKind"}) + if not (drops or adds_here): + continue + kept = [ + ast.unparse(alias) + for alias in statement.names + if (alias.asname or alias.name) not in fieldoffset_names + ] + text = ( + f"from {'.' * statement.level}{statement.module or ''} import {', '.join(kept)}\n" + if kept + else "" + ) + if adds_here: + text += f"from gt4py.next import {', '.join(missing)}\n" + added = True + _replace(module, statement, text) + if missing and not added: + module.notes.append( + f"{module.path}: import {', '.join(missing)} from 'gt4py.next', used by the migrated" + " declarations." + ) + + class _CartesianUses(ast.NodeVisitor): """Rewrite the uses of removed Cartesian offsets, and note the ones it cannot.""" @@ -159,6 +227,25 @@ def __init__(self, module: Module, cartesian: dict[str, str]) -> None: #: (line, start column, end column, replacement) of single-line expression rewrites self.rewrites: list[tuple[int, int, int, str]] = [] self.handled: set[int] = set() + #: names bound in the enclosing function scopes, which shadow a removed offset + self.shadowed: list[set[str]] = [] + + def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> None: + arguments = node.args + bound = { + argument.arg + for argument in (*arguments.posonlyargs, *arguments.args, *arguments.kwonlyargs) + } | { + name.id + for name in ast.walk(node) + if isinstance(name, ast.Name) and isinstance(name.ctx, ast.Store) + } + self.shadowed.append(bound) + self.generic_visit(node) + self.shadowed.pop() + + visit_AsyncFunctionDef = visit_FunctionDef + visit_Lambda = visit_FunctionDef def _rewrite(self, node: ast.expr, text: str) -> None: assert node.end_lineno is not None and node.end_col_offset is not None @@ -170,7 +257,9 @@ def _rewrite(self, node: ast.expr, text: str) -> None: def _dimension_of(self, node: ast.expr) -> str | None: """The dimension replacing `Off` or `module.Off`, qualified like the offset was.""" match node: - case ast.Name(id=name) if name in self.cartesian: + case ast.Name(id=name) if name in self.cartesian and not any( + name in bound for bound in self.shadowed + ): return self.cartesian[name] case ast.Attribute(value=value, attr=name) if name in self.cartesian: return f"{self.module.segment(value)}.{self.cartesian[name]}" @@ -225,6 +314,11 @@ def _migrate_cartesian_uses(module: Module, cartesian: dict[str, str]) -> None: for statement in module.tree.body: if not isinstance(statement, (ast.Import, ast.ImportFrom)): visitor.visit(statement) + match statement: + case ast.Assign(targets=[ast.Name(id="__all__")], value=ast.List(elts=all_names)): + for entry in all_names: + if isinstance(entry, ast.Constant) and entry.value in cartesian: + module.note(entry, f"'__all__' lists the removed offset '{entry.value}'.") lines = module.lines for line, start, end, text in sorted(visitor.rewrites, reverse=True): @@ -237,17 +331,32 @@ def _migrate_cartesian_uses(module: Module, cartesian: dict[str, str]) -> None: _PROVIDER_KEY_RE = re.compile(r"""(?P["'])(?P[A-Za-z_]\w*)(?P=quote)\s*:""") -def _report(module: Module, offset_names: set[str], dimension_names: set[str]) -> None: +def _report( + module: Module, + offset_keys: dict[str, str], + cartesian: dict[str, str], + dimension_names: set[str], +) -> None: for number, line in enumerate(module.source.splitlines(), start=1): for match in _PROVIDER_KEY_RE.finditer(line): - if match["name"] in offset_names: - module.notes.append( - f"{module.path}:{number}: offset-provider key '{match['name']}' is keyed by" - f" the connectivity class now, e.g. '{{{match['name']}: table}}'." + if (connectivity := offset_keys.get(match["name"])) is None: + continue + if connectivity in cartesian: + message = ( + f"offset-provider key '{match['name']}': remove the entry, a Cartesian shift" + f" ('{cartesian[connectivity]} + i') needs none." + ) + else: + message = ( + f"offset-provider key '{match['name']}' is keyed by the connectivity class" + f" now, e.g. '{{{connectivity}: table}}'." ) + module.notes.append(f"{module.path}:{number}: {message}") for node in ast.walk(module.tree): match node: - case ast.Attribute(value=ast.Name(id=name), attr="value") if name in dimension_names: + case ast.Attribute( + value=ast.Name(id=name) | ast.Attribute(attr=name), attr="value" + ) if name in dimension_names: module.note(node, f"'{name}.value': a dimension's name is '{name}.tag' now.") case ast.Call( func=ast.Name(id="isinstance"), args=[_, ast.Attribute(attr="Dimension")] @@ -272,17 +381,24 @@ def migrate(sources: dict[pathlib.Path, str]) -> tuple[dict[pathlib.Path, str], Module(path=path, source=source, tree=ast.parse(source)) for path, source in sources.items() ] cartesian: dict[str, str] = {} - offset_names: set[str] = set() + #: offset-provider keys that name a `FieldOffset`: its variable name and its tag, if different + offset_keys: dict[str, str] = {} dimension_names: set[str] = set() for module in modules: - for _, name, _, _, kind_of_call in _declarations(module): - (dimension_names if kind_of_call == "Dimension" else offset_names).add(name) + for _, name, call, _, kind_of_call in _declarations(module): + if kind_of_call == "Dimension": + dimension_names.add(name) + continue + offset_keys[name] = name + if call.args and isinstance(tag := call.args[0], ast.Constant): + offset_keys[str(tag.value)] = name _migrate_declarations(module, cartesian) + _migrate_imports(module) results: dict[pathlib.Path, str] = {} notes: list[str] = [] for module in modules: - _report(module, offset_names, dimension_names) + _report(module, offset_keys, cartesian, dimension_names) migrated = _apply(module) # Cartesian uses are rewritten on the migrated text, re-parsed, so that line numbers # refer to what the declaration edits left. @@ -327,7 +443,7 @@ def run( ) ) if notes: - typer.echo("\nLeft to migrate by hand:", err=True) + typer.echo("\nLeft to migrate by hand (line numbers of the original files):", err=True) for note in notes: typer.echo(f" {note}", err=True) diff --git a/scripts/tests/python/test_migrate_connectivities.py b/scripts/tests/python/test_migrate_connectivities.py index a8d48ad26d..abcb777af3 100644 --- a/scripts/tests/python/test_migrate_connectivities.py +++ b/scripts/tests/python/test_migrate_connectivities.py @@ -114,3 +114,59 @@ def test_what_is_left_is_reported(): assert any("offset-provider key 'E2C'" in note for note in notes) assert any("offset-provider key 'Koff'" in note for note in notes) assert any("'KDim.value'" in note for note in notes) + + +BARE = textwrap.dedent( + """\ + from gt4py.next import Dimension, DimensionKind, FieldOffset as FO + + LOCAL = DimensionKind.LOCAL + Vertex = Dimension("Vertex") + Edge = Dimension("Edge") + V2EDim = Dimension("V2E", LOCAL) + V2E = FO("V2E_TAG", source=Edge, target=(Vertex, V2EDim)) + """ +) + + +def test_unqualified_names_and_aliases(): + results, notes = _migrate(bare=BARE, user='table = {"V2E_TAG": t}\n') + migrated = results["bare"] + + assert "FO" not in migrated + assert "from gt4py.next import Dimension, DimensionKind\n" in migrated + assert ( + "from gt4py.next import DimensionIndex, LocalDimensionIndex, NeighborConnectivity\n" + in migrated + ) + assert "class V2EDim(LocalDimensionIndex): ..." in migrated + assert "class V2E(NeighborConnectivity[Vertex, Edge]):" in migrated + namespace: dict = {"__name__": "migrated_bare"} + exec(migrated, namespace) + assert namespace["V2E"].Local is namespace["V2EDim"] + # a key spelled with the offset's tag is reported, naming the class + assert any("'V2E_TAG'" in note and "{V2E: table}" in note for note in notes) + + +def test_shadowed_names_and_all_are_left_alone(): + source = textwrap.dedent( + """\ + from pkg.dimension import KDim, Koff + + __all__ = ["Koff"] + + + def helper(Koff): + return Koff + 1 + + + def uses(a): + return a(Koff[1]) + a(dims.KDim.value) + """ + ) + results, notes = _migrate(dimension=DIMENSIONS, user=source) + + assert "return Koff + 1" in results["user"] + assert "a(KDim + 1)" in results["user"] + assert any("'__all__' lists the removed offset 'Koff'" in note for note in notes) + assert any("'KDim.value'" in note for note in notes) diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 3128b6e535..06da46e380 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -352,7 +352,6 @@ def staggered_base_tag(tag: Tag) -> Optional[Tag]: return match["base"] if (match := _STAGGERED_TAG_RE.match(tag)) is not None else None -@functools.cache def resolve(tag: Tag) -> Dimension: """ Return the dimension class a tag names, by importing it. @@ -365,6 +364,10 @@ def resolve(tag: Tag) -> Dimension: so the longest importable prefix wins and the remainder is walked as attributes. A collision would need a module path and an attribute chain to have the same spelling. + Only where the module path ends is memoized; the attribute walk is repeated on every call, so + a declaration redefined under the same name (e.g. by re-running a notebook cell) resolves to + the new class. + Parametrized dimensions such as `Staggered[K]` have no importable qualname; their tag has the form `[]` and is resolved by subscripting the owner, which goes through its intern table and so returns the identical class. @@ -396,26 +399,33 @@ def resolve(tag: Tag) -> Dimension: return cast(Dimension, obj) -@functools.cache def _import_qualified_name(tag: Tag) -> Any: """Import the object a dotted qualified name refers to; see `resolve`.""" + module_name, attrs = _split_qualified_name(tag) + obj: Any = sys.modules.get(module_name) or importlib.import_module(module_name) + for attr in attrs: + try: + obj = getattr(obj, attr) + except AttributeError as ex: + raise ValueError( + f"Cannot resolve tag '{tag}': '{module_name}' has no attribute '{'.'.join(attrs)}'." + ) from ex + return obj + + +@functools.cache +def _split_qualified_name(tag: Tag) -> tuple[str, tuple[str, ...]]: + """Split a dotted name at its longest importable module prefix: `(module, attributes)`.""" parts = tag.split(".") for split in range(len(parts), 0, -1): + module_name = ".".join(parts[:split]) try: - obj: Any = importlib.import_module(".".join(parts[:split])) + importlib.import_module(module_name) except ImportError: continue - for attr in parts[split:]: - try: - obj = getattr(obj, attr) - except AttributeError as ex: - raise ValueError( - f"Cannot resolve dimension tag '{tag}': '{'.'.join(parts[:split])}' has" - f" no attribute '{attr}'." - ) from ex - return obj + return module_name, tuple(parts[split:]) raise ValueError( - f"Cannot resolve dimension tag '{tag}': no importable module prefix. A dimension" + f"Cannot resolve tag '{tag}': no importable module prefix. A dimension or connectivity" " referenced from the IR must be declared at module level in an importable module." ) @@ -2284,14 +2294,27 @@ def fail(reason: str) -> NoReturn: if not isinstance(table_type, NeighborConnectivityType): fail(f"expected a neighbor table, got '{table_type}'") + + def redefined(found: Sequence[Dimension], expected: Sequence[Dimension]) -> str: + if any(f is not e and f.tag == e.tag for f, e in zip(found, expected)): + return ( + " (a dimension of the same name but a different class: was the declaration" + " redefined, e.g. by re-running a notebook cell?)" + ) + return "" + 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))})'" + + redefined(table_type.domain, expected_domain) ) if table_type.codomain is not connectivity.codomain: - fail(f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'") + fail( + f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'" + + redefined((table_type.codomain,), (connectivity.codomain,)) + ) 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: @@ -2332,7 +2355,8 @@ def as_tag_keyed_offset_provider( """ Key an offset provider by tags, the form the IR and the backends use. - A `NeighborConnectivity` key becomes its local dimension's tag. A string key is taken to be + A `NeighborConnectivity` key becomes its `offset_tag`: its local dimension's tag, or its own + for a connectivity sharing another one's local dimension. A string key is taken to be such a tag already, and is rejected if it cannot be one: a tag is a qualified name, so a bare name such as `"V2E"` is the removed `FieldOffset` spelling. @@ -2428,10 +2452,12 @@ def _check_shared_local_dimensions( and is_neighbor_table(first) and is_neighbor_table(table) ): - same_structure = bool( - np.array_equal( - first.asnumpy() == first_type.skip_value, - table.asnumpy() == table_type.skip_value, + # NOTE: compared where the tables live, without copying device arrays to the host. + xp = first.array_ns # type: ignore[attr-defined] # all tables are NdArrayFields + same_structure = first.ndarray.shape == table.ndarray.shape and bool( + xp.all( + (first.ndarray == first_type.skip_value) + == (xp.asarray(table.ndarray) == table_type.skip_value) ) ) if not same_structure: diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index 37a14db1e1..e226d4e244 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -729,6 +729,10 @@ class FieldOperatorFromFoast(FieldOperator): @override def __call__(self, *args: Any, **kwargs: Any) -> Any: assert self.backend is not None + if "offset_provider" in kwargs: + kwargs["offset_provider"] = common.as_tag_keyed_offset_provider( + kwargs["offset_provider"] + ) compiled_fo = self.backend.compile( self.foast_stage, arguments.CompileTimeArgs.from_concrete(*args, **kwargs) ) diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index f5bb207efb..cc5ab3642e 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -41,7 +41,8 @@ ScalarOrTupleOfScalars: TypeAlias = xtyping.MaybeNestedInTuple[core_defs.Scalar] -#: Content of the key: (*hashable_arg_descriptors, id(offset_provider), concrete_instantation_if_generic) +#: Content of the key: (*hashable_arg_descriptors, hash of the offset provider's (tag, id(table)) +#: items, concrete_instantation_if_generic) CompiledProgramsKey: TypeAlias = tuple[tuple[Hashable, ...], int, str | None] ArgStaticDescriptorsByType: TypeAlias = dict[ diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py index 1af014470c..04ff03a054 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py @@ -92,14 +92,14 @@ def _get_args(sdfg: dace.SDFG, args: Sequence[Any]) -> dict[str, Any]: def get_sdfg_conn_args( sdfg: dace.SDFG, - offset_provider: gtx_common.OffsetProvider, + offset_provider: gtx_common.OffsetProviderLike, ) -> dict[str, core_defs.NDArrayObject]: """ Extracts the connectivity tables that are used in the sdfg and ensures that the memory buffers are allocated for the target device. """ connectivity_args = {} - for offset, connectivity in offset_provider.items(): + for offset, connectivity in gtx_common.as_tag_keyed_offset_provider(offset_provider).items(): name = gtx_dace_args.connectivity_identifier(offset) if name in sdfg.arrays: assert gtx_common.is_neighbor_table(connectivity) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 69072d70b2..0f3a549160 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -596,6 +596,38 @@ def test_nothing_bound(self): common.connectivity_key_over({E2V.offset_tag: self._type(E2V)}, V2E.Local) +def test_redefined_declaration_resolves_to_the_new_class(monkeypatch): + """Re-running a notebook cell redefines declarations under the same names.""" + import sys + import types as pytypes + + module = pytypes.ModuleType("_redefined_connectivity_module") + monkeypatch.setitem(sys.modules, module.__name__, module) + source = textwrap.dedent( + """ + from gt4py.next.common import DimensionIndex, LocalDimensionIndex, NeighborConnectivity + + class V(DimensionIndex): ... + class E(DimensionIndex): ... + class V2E(NeighborConnectivity[V, E], max_neighbors={n}): + class Local(LocalDimensionIndex): ... + """ + ) + exec(source.format(n=4), module.__dict__) + old = module.V2E + common.check_offset_provider({old: _table(domain=(module.V, old.Local), codomain=module.E)}) + + exec(source.format(n=2), module.__dict__) + new = module.V2E + assert common.resolve(new.Local.tag) is new.Local + common.check_offset_provider( + {new: _table(domain=(module.V, new.Local), codomain=module.E, data=((0, 1), (1, 0)))} + ) + with pytest.raises(ValueError, match="was the declaration redefined"): + common.check_neighbor_table( + new, _table(domain=(old.origin, old.Local), codomain=module.E, data=((0, 1), (1, 0))) + + def test_the_const_list_dimension_cannot_be_adopted(): with pytest.raises(TypeError, match="cannot adopt"): _declare( From f80624f414a14d1f5a79f4b3755b91d4d9d113f1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 16:10:06 +0200 Subject: [PATCH 12/17] fix[next]: review fixes for class-keyed offset providers - check the provider at the entry points that only normalized it: the iterator 'fendef', 'FieldOperatorFromFoast' and the DaCe orchestration's connectivities - remember a checked provider by the identity of its tables, and read the tables (comparing skip positions of shared local dimensions) only where a program is compiled, not on the call path - reject a key that names the connectivity rather than the declaration - cache a key that is not a qualified name, so it costs one import attempt - DaCe 'get_sdfg_conn_args' is an IR-level hook, so it is not strict either - the migration script writes 'Local: TypeAlias = ...' and imports 'typing' - PR 1's bare-Cartesian-offset check is unreachable without 'FieldOffset' - ADR 0029: which entry points are strict, and what the checks read --- .../ADRs/next/0029-Connectivities_As_Types.md | 38 ++++---- scripts/python/migrate_connectivities.py | 27 +++++- .../python/test_migrate_connectivities.py | 8 +- src/gt4py/next/common.py | 91 +++++++++++++++---- src/gt4py/next/ffront/decorator.py | 1 + .../ffront/foast_passes/type_deduction.py | 19 +--- src/gt4py/next/iterator/runtime.py | 1 + src/gt4py/next/otf/compiled_program.py | 2 +- src/gt4py/next/otf/options.py | 2 + .../runners/dace/sdfg_callable.py | 5 +- .../ffront_tests/test_diagnostic_messages.py | 16 ---- .../unit_tests/test_neighbor_connectivity.py | 36 ++++++-- 12 files changed, 164 insertions(+), 82 deletions(-) diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index 0bbc6bf82b..dec623ac50 100644 --- a/docs/development/ADRs/next/0029-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -130,11 +130,10 @@ whose tag is the connectivity's `offset_tag`: 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 +`V2E.Local` inside DSL code types as that local dimension. The other frontend +touch points treat a declaration as the offset it replaces: 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. ### Offset providers are keyed by the declaration @@ -145,18 +144,23 @@ Users bind tables to declarations: program(..., offset_provider={V2E: v2e_table, C2E: c2e_table}) ``` -Every entry point of a program (`Program.__call__`, `FieldOperator.__call__`, -`compile`, `CompilationOptions.connectivities`, `embedded.context.update`, the -iterator `fendef`) normalizes such a provider to the form the IR uses: each -declaration is replaced by its `offset_tag`. Everything below the entry points — -lowering, the backends, compiled-program caching — therefore keeps seeing a -provider keyed by strings, which is also what hand-written IR uses. A string key -must be a tag, i.e. a qualified name; a bare name such as `"V2E"` is the removed -`FieldOffset` spelling and is rejected with a message pointing here. - -Tables are checked against their declarations (`check_offset_provider`) once per -compiled variant and on each embedded call, not on every compiled call: the -check builds the table's type, which is too slow for the call path. A tag that +Every entry point of a program normalizes such a provider to the form the IR +uses: each declaration is replaced by its `offset_tag`. Everything below the +entry points — lowering, the backends, compiled-program caching — therefore keeps +seeing a provider keyed by strings, which is also what hand-written IR uses. + +The frontend entry points (`Program.__call__`, `FieldOperator.__call__`, +`compile`, `CompilationOptions.connectivities`) are *strict*: a string key must +be a tag, i.e. a qualified name, and a bare name such as `"V2E"` is the removed +`FieldOffset` spelling, rejected with a message pointing here. The IR-level hooks +(`embedded.context.update`, the iterator `fendef`, DaCe's `get_sdfg_conn_args`) +accept any string, because a hand-written program names its offsets itself. + +Tables are checked against their declarations (`check_offset_provider`) at every +entry point, but the result is remembered per set of bound tables, so repeated +calls cost one hash. Reading the tables — comparing the skip-value positions of +two connectivities that share a local dimension — is done only where a program is +compiled, not on the call path. A tag that names no declared connectivity, as in hand-written IR, is not checked. ### `FieldOffset` is removed diff --git a/scripts/python/migrate_connectivities.py b/scripts/python/migrate_connectivities.py index 304f7401b2..d3523b1887 100644 --- a/scripts/python/migrate_connectivities.py +++ b/scripts/python/migrate_connectivities.py @@ -24,10 +24,10 @@ class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... class E2CDim(gtx.LocalDimensionIndex): ... class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): - Local = E2CDim + Local: typing.TypeAlias = E2CDim ... a(KDim + 1) ... as_offset(KDim, k_field) ... -A connectivity adopts its existing local dimension (`Local = E2CDim`), so the names already used +A connectivity adopts its existing local dimension (`Local: TypeAlias = E2CDim`), so the names already used for local dimensions keep working, and so do offsets that share a local dimension (`C2CE` with `C2EDim`). What cannot be rewritten from the source alone is reported instead: offset providers keyed by strings, which become keyed by the connectivity (`{E2C: table}`), `.value` @@ -69,6 +69,8 @@ class Module: notes: list[str] = dataclasses.field(default_factory=list) #: Names of `gt4py.next` the migrated declarations use unqualified, to be imported. needed: set[str] = dataclasses.field(default_factory=set) + #: Whether the migrated declarations need `typing` imported (for `Local: TypeAlias = ...`). + needed_typing: bool = False @property def lines(self) -> list[str]: @@ -163,8 +165,11 @@ def _migrate_declarations(module: Module, cartesian: dict[str, str]) -> None: origin, local = (module.segment(element) for element in target.elts) text = ( f"class {name}({prefix}NeighborConnectivity[{origin}, {module.segment(source)}]):\n" - f" Local = {local}\n" + # NOTE: `TypeAlias`, not a plain assignment: it is what keeps the adopted local + # dimension a *type* for mypy (see ADR 0029). + f" Local: typing.TypeAlias = {local}\n" ) + module.needed_typing = True if not prefix: module.needed.add("NeighborConnectivity") _replace(module, statement, text) @@ -188,6 +193,10 @@ def _migrate_imports(module: Module) -> None: for alias in statement.names } missing = sorted(module.needed - imported) + typing_missing = module.needed_typing and not any( + isinstance(statement, ast.Import) and any(a.name == "typing" for a in statement.names) + for statement in ast.walk(module.tree) + ) added = False for statement in module.tree.body: if not isinstance(statement, ast.ImportFrom): @@ -211,6 +220,18 @@ def _migrate_imports(module: Module) -> None: text += f"from gt4py.next import {', '.join(missing)}\n" added = True _replace(module, statement, text) + if typing_missing: + # `Local: TypeAlias = ...` needs it; the first import statement is a safe place + for statement in module.tree.body: + if isinstance(statement, (ast.Import, ast.ImportFrom)): + module.edits.append( + Edit(statement.lineno - 1, statement.lineno - 1, "import typing\n") + ) + break + else: + module.notes.append( + f"{module.path}: import 'typing', used by the migrated declarations." + ) if missing and not added: module.notes.append( f"{module.path}: import {', '.join(missing)} from 'gt4py.next', used by the migrated" diff --git a/scripts/tests/python/test_migrate_connectivities.py b/scripts/tests/python/test_migrate_connectivities.py index abcb777af3..4d196eeeba 100644 --- a/scripts/tests/python/test_migrate_connectivities.py +++ b/scripts/tests/python/test_migrate_connectivities.py @@ -68,6 +68,7 @@ def test_declarations(): assert results["dimension"] == textwrap.dedent( """\ + import typing import gt4py.next as gtx class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @@ -77,11 +78,11 @@ class CEDim(gtx.DimensionIndex): ... class E2CDim(gtx.LocalDimensionIndex): ... class C2EDim(gtx.LocalDimensionIndex): ... class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): - Local = E2CDim + Local: typing.TypeAlias = E2CDim class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): - Local = C2EDim + Local: typing.TypeAlias = C2EDim class C2CE(gtx.NeighborConnectivity[CellDim, CEDim]): - Local = C2EDim + Local: typing.TypeAlias = C2EDim """ ) @@ -141,6 +142,7 @@ def test_unqualified_names_and_aliases(): ) assert "class V2EDim(LocalDimensionIndex): ..." in migrated assert "class V2E(NeighborConnectivity[Vertex, Edge]):" in migrated + assert " Local: typing.TypeAlias = V2EDim" in migrated namespace: dict = {"__name__": "migrated_bare"} exec(migrated, namespace) assert namespace["V2E"].Local is namespace["V2EDim"] diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 06da46e380..f3b6282e48 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -399,6 +399,18 @@ def resolve(tag: Tag) -> Dimension: return cast(Dimension, obj) +def _import_qualified_name_or_none(tag: Tag) -> Any: + """`_import_qualified_name`, returning `None` for a name that does not resolve.""" + if (split := _split_qualified_name_or_none(tag)) is None: + return None + module_name, attrs = split + obj: Any = sys.modules.get(module_name) or importlib.import_module(module_name) + for attr in attrs: + if (obj := getattr(obj, attr, None)) is None: + return None + return obj + + def _import_qualified_name(tag: Tag) -> Any: """Import the object a dotted qualified name refers to; see `resolve`.""" module_name, attrs = _split_qualified_name(tag) @@ -413,6 +425,21 @@ def _import_qualified_name(tag: Tag) -> Any: return obj +@functools.cache +def _split_qualified_name_or_none(tag: Tag) -> Optional[tuple[str, tuple[str, ...]]]: + """ + `_split_qualified_name`, returning `None` instead of raising. + + Separate and memoized so that a string that is not a qualified name -- an offset-provider key + of hand-written IR, say -- costs one import attempt in total, not one per call. A module whose + import *fails* other than by not being found is reported, not cached away. + """ + try: + return _split_qualified_name(tag) + except ValueError: + return None + + @functools.cache def _split_qualified_name(tag: Tag) -> tuple[str, tuple[str, ...]]: """Split a dotted name at its longest importable module prefix: `(module, attributes)`.""" @@ -1549,7 +1576,9 @@ def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: return True -def hash_offset_provider_items_by_id(offset_provider: OffsetProvider) -> int: +def hash_offset_provider_items_by_id( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, +) -> int: """ Compute hash of an offset provider on the tuples of key and value id. @@ -2121,7 +2150,10 @@ def bound_table(cls) -> NeighborTable: f"'{cls.__qualname__}' can only be resolved to a table during embedded execution." ) table = get_offset(offset_provider, cls.offset_tag) - assert is_neighbor_table(table) + if not is_neighbor_table(table): + raise TypeError( + f"'{cls.__qualname__}' is bound to '{table}', which is not a neighbor table." + ) return table @@ -2387,42 +2419,66 @@ def _check_tag_keys(offset_provider: Mapping[Any, Any]) -> None: for key in offset_provider: if not isinstance(key, str) or "." not in key: raise TypeError( - f"Invalid offset-provider key '{key!r}': offset providers are keyed by" + f"Invalid offset-provider key {key!r}: offset providers are keyed by" " 'NeighborConnectivity' declarations, e.g. '{V2E: v2e_table}'. A bare name is the" " spelling of the removed 'FieldOffset' (see ADR 0029)." ) -def check_offset_provider(offset_provider: OffsetProviderLike | OffsetProviderTypeLike) -> None: +#: Offset providers already checked, by the hash of their `(key, id(table))` items. Bounded, and +#: not authoritative: like the compiled-program cache (which keys on the same hash), it can in +#: principle skip a check when a freed table is replaced at the same address. See +#: `check_offset_provider`. +_CHECKED_OFFSET_PROVIDERS: Final[collections.OrderedDict[int, None]] = collections.OrderedDict() +_CHECKED_OFFSET_PROVIDERS_MAX: Final = 256 + + +def check_offset_provider( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, *, deep: bool = False +) -> None: """ - Check every table of a tag-keyed offset provider against its connectivity declaration. + Check every table of an offset provider against its connectivity declaration. - A tag that does not name the local dimension of a declared connectivity -- e.g. one used only - by hand-written IR -- has no declaration to be checked against and is skipped. + A key that does not name a declared connectivity -- e.g. a tag used only by hand-written IR -- + has no declaration to be checked against and is skipped. Providers are remembered by the + identity of their tables, so repeated calls with the same tables cost one hash. + + Args: + offset_provider: The provider, keyed by declarations or by tags. + deep: Also compare the skip-value positions of tables over one shared local dimension, + which reads the tables. The compile path does; the call path does not. Raises: ValueError: If a table does not match its declaration, see `check_neighbor_table`. """ + if (seen := hash_offset_provider_items_by_id(offset_provider)) in _CHECKED_OFFSET_PROVIDERS: + return for key, table in offset_provider.items(): declaration: Any = key if isinstance(key, str): - try: - declaration = _import_qualified_name(key) - except ValueError: + if (declaration := _import_qualified_name_or_none(key)) is None: continue if isinstance(declaration, DimensionMeta): # the local dimension's tag names its owner's table declaration = getattr(declaration, "owner", None) - if isinstance(declaration, ConnectivityMeta) and declaration.offset_tag in ( - key, - getattr(key, "offset_tag", None), - ): + if isinstance(declaration, ConnectivityMeta): + if declaration.offset_tag not in (key, getattr(key, "offset_tag", None)): + # e.g. `{V2E.tag: table}`: the connectivity's own tag, which is the IR name only + # of a connectivity *sharing* a local dimension + raise ValueError( + f"Invalid offset-provider key '{key}': it names the connectivity" + f" '{declaration.__qualname__}', whose key is the declaration itself" + f" ('{{{declaration.__qualname__}: table}}')." + ) check_neighbor_table(cast(type[NeighborConnectivity], declaration), table) - _check_shared_local_dimensions(offset_provider) + _check_shared_local_dimensions(offset_provider, deep=deep) + _CHECKED_OFFSET_PROVIDERS[seen] = None + while len(_CHECKED_OFFSET_PROVIDERS) > _CHECKED_OFFSET_PROVIDERS_MAX: + _CHECKED_OFFSET_PROVIDERS.popitem(last=False) def _check_shared_local_dimensions( - offset_provider: OffsetProviderLike | OffsetProviderTypeLike, + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, *, deep: bool = False ) -> None: """ Check that the tables over one local dimension have the same neighbor structure. @@ -2447,7 +2503,8 @@ def _check_shared_local_dimensions( first_type.has_skip_values, ) if ( - same_structure + deep + and same_structure and first_type.has_skip_values and is_neighbor_table(first) and is_neighbor_table(table) diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index e226d4e244..d409428b24 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -733,6 +733,7 @@ def __call__(self, *args: Any, **kwargs: Any) -> Any: kwargs["offset_provider"] = common.as_tag_keyed_offset_provider( kwargs["offset_provider"] ) + common.check_offset_provider(kwargs["offset_provider"]) compiled_fo = self.backend.compile( self.foast_stage, arguments.CompileTimeArgs.from_concrete(*args, **kwargs) ) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index bbdf8433fe..9aa62ecae7 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -793,20 +793,11 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call: ): raise errors.DSLError(node.location, "Functions can only be called directly.") elif isinstance(new_func.type, ts.FieldType): - for arg in new_args: - # A Cartesian `FieldOffset` shifts by the index it is subscripted with, so it - # carries no displacement on its own. Only an offset with a local dimension is - # 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 - ): - raise errors.DSLError( - arg.location, - f"Cannot shift by the Cartesian offset '{arg!s}' without an index.", - hints=[f"Give the displacement, e.g. '{arg!s}[1]'."], - ) + # NOTE: a bare single-target offset used to be rejected here, as the unsubscripted + # Cartesian `FieldOffset` `a(Koff)`. There is no such declaration any more: a Cartesian + # shift is `a(Dim + i)`, and the only single-target offset left is the result of + # `as_offset`, which is a call, not a name. + pass elif isinstance(new_func.type, ts.DimensionType): assert new_func.type.dim.kind == DimensionKind.LOCAL return foast.Call( diff --git a/src/gt4py/next/iterator/runtime.py b/src/gt4py/next/iterator/runtime.py index 9df3f64e1a..2df378346c 100644 --- a/src/gt4py/next/iterator/runtime.py +++ b/src/gt4py/next/iterator/runtime.py @@ -82,6 +82,7 @@ def __call__( offset_provider = common.as_tag_keyed_offset_provider( offset_provider or self.offset_provider, strict=False ) + common.check_offset_provider(offset_provider) column_axis = column_axis or self.column_axis if backend is not None: diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index cc5ab3642e..6852a27619 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -644,7 +644,7 @@ def _compile_variant( else: raise ValueError(f"Invalid 'offset_provider': {offset_provider}") - common.check_offset_provider(offset_provider) + common.check_offset_provider(offset_provider, deep=True) self._initialize_argument_descriptor_mapping(argument_descriptors) _validate_argument_descriptors(self.program_type, argument_descriptors) diff --git a/src/gt4py/next/otf/options.py b/src/gt4py/next/otf/options.py index c368e524f7..1056e2da5d 100644 --- a/src/gt4py/next/otf/options.py +++ b/src/gt4py/next/otf/options.py @@ -45,6 +45,8 @@ def __post_init__(self) -> None: object.__setattr__( self, "connectivities", common.as_tag_keyed_offset_provider(self.connectivities) ) + # the DaCe orchestration reads these directly, without passing an offset provider + common.check_offset_provider(self.connectivities, deep=True) assert CompilationOptionsArgs.__annotations__.keys() == CompilationOptions.__annotations__.keys() diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py index 04ff03a054..d21c475723 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py @@ -99,7 +99,10 @@ def get_sdfg_conn_args( that the memory buffers are allocated for the target device. """ connectivity_args = {} - for offset, connectivity in gtx_common.as_tag_keyed_offset_provider(offset_provider).items(): + # NOTE: not strict, like the other IR-level hooks: the keys of a hand-written program are its + # own business, and a declaration is normalized to its tag either way. + provider = gtx_common.as_tag_keyed_offset_provider(offset_provider, strict=False) + for offset, connectivity in provider.items(): name = gtx_dace_args.connectivity_identifier(offset) if name in sdfg.arrays: assert gtx_common.is_neighbor_table(connectivity) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py index 337c7ce67f..1fc3ca6782 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py @@ -30,8 +30,6 @@ class IDim(gtx.DimensionIndex): ... -IOff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) - # A PEP 695 alias whose value raises when it is evaluated, standing in for the # common case of a typo'd dtype ('np.foat64') inside an alias definition. _empty_module = types.ModuleType("_empty_module") @@ -361,20 +359,6 @@ def broken(a: BrokenFieldAlias) -> gtx.Field[[IDim], float64]: assert re.search(r"\| +\^{19}(?!\^)", str(err)), str(err) -def test_unindexed_cartesian_offset_names_the_offset_as_written(): - # The tag of 'IOff' is 'Ioff'; the message has to quote what the user wrote. - def unindexed(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: - return a(IOff) - - err = parse_error(unindexed) - - assert err.message == "Cannot shift by the Cartesian offset 'IOff' without an index." - assert err.hints == ["Give the displacement, e.g. 'IOff[1]'."] - rendered = str(err) - assert "return a(IOff)" in rendered - assert re.search(r"\| +\^{4}(?!\^)", rendered), rendered - - def test_indexed_dimension_shift_is_rejected_with_a_hint(): def indexed_shift(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: return a((IDim + 1)[0]) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 0f3a549160..2aaa12c46e 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -452,6 +452,14 @@ def test_grid_type_deduction(self): def _table(domain=(Vertex, V2E.Local), codomain=Edge, data=((0, 1, 2, 3), (1, 2, 3, 0))): from gt4py.next import constructors + data = np.array(data) + return constructors.as_connectivity( + domain=dict(zip(domain, data.shape)), + codomain=codomain, + data=data, + skip_value=common._DEFAULT_SKIP_VALUE, + ) + def test_redefined_declaration_with_an_adopted_local(monkeypatch): """Re-running a cell must re-own the adopted local dimension, not become a sharer.""" @@ -495,14 +503,6 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): with pytest.raises(TypeError, match="not a connectivity declaration"): common.local_dimension_of(NeighborConnectivity) - data = np.array(data) - return constructors.as_connectivity( - domain=dict(zip(domain, data.shape)), - codomain=codomain, - data=data, - skip_value=common._DEFAULT_SKIP_VALUE, - ) - class TestOffsetProvider: def test_class_keys_become_tags(self): @@ -544,9 +544,24 @@ def test_sharing_connectivities_need_the_same_skip_positions(self): owner = _table(data=((0, 1, 2, 3), (1, 2, 3, 0))) consistent = _table(data=((3, 2, 1, 0), (0, 3, 2, 1))) inconsistent = _table(data=((3, 2, 1, -1), (0, 3, 2, 1))) - common.check_offset_provider({V2E: owner, V2EShared: consistent}) + common.check_offset_provider({V2E: owner, V2EShared: consistent}, deep=True) with pytest.raises(ValueError, match="different neighbor structure"): - common.check_offset_provider({V2E: owner, V2EShared: inconsistent}) + common.check_offset_provider({V2E: owner, V2EShared: inconsistent}, deep=True) + # the call path does not read the tables + common.check_offset_provider({V2E: owner, V2EShared: inconsistent}) + + def test_the_connectivity_s_own_tag_is_rejected(self): + with pytest.raises(ValueError, match="whose key is the declaration itself"): + common.check_offset_provider({V2E.tag: _table()}) + + def test_a_checked_provider_is_remembered(self): + provider = {V2E: _table(codomain=Vertex)} + with pytest.raises(ValueError, match="does not match its declaration"): + common.check_offset_provider(provider) + # and a provider that passes is not checked twice + good = {V2E: _table()} + common.check_offset_provider(good) + common.check_offset_provider(good) def test_check_skips_undeclared_tags(self): common.check_offset_provider({"some.hand.written.tag": _table()}) @@ -626,6 +641,7 @@ class Local(LocalDimensionIndex): ... with pytest.raises(ValueError, match="was the declaration redefined"): common.check_neighbor_table( new, _table(domain=(old.origin, old.Local), codomain=module.E, data=((0, 1), (1, 0))) + ) def test_the_const_list_dimension_cannot_be_adopted(): From 0619aca8afbf5214c5220bb50c66adb8ee9176f1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 16:20:12 +0200 Subject: [PATCH 13/17] docs[next]: ts.OffsetType stays, with the reason where the TODO was --- src/gt4py/next/type_system/type_specifications.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 50f451fd39..f5e0ea770c 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -71,7 +71,9 @@ def __str__(self) -> str: class OffsetType(TypeSpec): - # TODO(havogt): replace by ConnectivityType + # NOTE: kept, against the TODO that stood here: since ADR 0029 this types a connectivity + # declaration (`V2E.__gt_type__()`) and the result of `as_offset`, and a `ConnectivityType` + # is what the *bound table* produces. Renaming it would be churn with no user-visible gain. source: common.Dimension target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] #: The offset-provider key; `None` for the untagged Cartesian `Dim + offset`. From e680735de26974e16ee7fafef0bcb062753b887b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 04:54:51 +0200 Subject: [PATCH 14/17] refactor[next]: ConstList replaces _CONST_DIM; AxisLiteral drops kind MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - common.ConstList: the owner-less, size-1 local dimension of make_const_list, used directly by iterator embedded and the DaCe lowering (identity checks) - AxisLiteral stores only the tag; kind and dim are resolved from it. The pretty printer derives the suffix (now with ₗ for local dimensions, which used to be printed as vertical), and the parser ignores it. --- .../next/0028-Dimensions_As_Nominal_Types.md | 4 ++- src/gt4py/next/common.py | 17 +++++------- src/gt4py/next/ffront/foast_to_gtir.py | 4 +-- src/gt4py/next/ffront/past_to_itir.py | 6 ++--- src/gt4py/next/iterator/embedded.py | 18 +++++-------- src/gt4py/next/iterator/ir.py | 15 ++++++++--- src/gt4py/next/iterator/ir_utils/ir_makers.py | 4 +-- src/gt4py/next/iterator/pretty_parser.py | 7 +++-- src/gt4py/next/iterator/pretty_printer.py | 16 ++++++++--- src/gt4py/next/iterator/tracing.py | 2 +- .../iterator/transforms/remove_broadcast.py | 4 +-- .../dace/lowering/gtir_to_sdfg_lambda.py | 24 +++++++---------- .../ffront_tests/test_foast_to_gtir.py | 4 +-- .../test_embedded_field_with_list.py | 5 ++-- .../iterator_tests/test_pretty_parser.py | 4 +-- .../iterator_tests/test_pretty_printer.py | 27 ++++++++++++------- .../iterator_tests/test_pretty_roundtrip.py | 17 +++--------- .../iterator_tests/test_type_inference.py | 20 +++++--------- 18 files changed, 94 insertions(+), 104 deletions(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 224f8ec289..295ceebe59 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -78,7 +78,9 @@ disappears. 1. **Reconstruction from the IR is an import.** `common.resolve(tag)` imports the module and walks the qualname; nested declarations resolve naturally. The IR references a Python type exactly the way `pickle` references a class. It is - memoized, because type inference calls it once per `AxisLiteral`. + memoized, because type inference calls it once per `AxisLiteral`. An + `AxisLiteral` stores only the tag: its `kind` is the resolved dimension's, so the + two cannot disagree. A purely dotted tag does not record where the module path ends and the qualname begins, so `resolve` tries the *longest importable prefix* and walks diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index f3b6282e48..ef43d5593c 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -2038,19 +2038,16 @@ def _check_neighbor_count(cls: type, name: str, count: Optional[int]) -> Optiona return int(count) -class ConstListDim(LocalDimensionIndex): +class ConstList(LocalDimensionIndex, size=1): """ - The local dimension of a list whose length is known at compile time (`make_const_list`). + The local dimension of a list of one repeated value (`make_const_list`). - Declared here, once, because it must be a *single* class. It used to be built - independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless - while dimensions compared by `(name, kind)` -- the two instances were equal. Under - nominal identity (ADR 0028) two declarations would be two different dimensions, and the - `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s - built by embedded execution. + 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. - TODO: becomes an owner-less local dimension with an explicit size, generalising this from - length 1 to length *n*, once local dimensions know their connectivity. + 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__ = () diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 1b5aabf37d..d1b0a6cbb1 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -234,12 +234,12 @@ def visit_Symbol(self, node: foast.Symbol, **kwargs: Any) -> itir.Sym: def visit_Name(self, node: foast.Name, **kwargs: Any) -> itir.SymRef | itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag) return im.ref(node.id) def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag) if isinstance(named_tup_type := node.value.type, ts.NamedCollectionType): ind = named_tup_type.keys.index(node.attr) diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index a00eccf544..57787f4d1f 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -382,9 +382,7 @@ def _construct_itir_domain_arg( domain_args = [] for dim_i, dim in enumerate(out_type.dims): # an expression for the range of a dimension - dim_range = im.call("get_domain_range")( - out_expr, itir.AxisLiteral(value=dim.tag, kind=dim.kind) - ) + dim_range = im.call("get_domain_range")(out_expr, itir.AxisLiteral(value=dim.tag)) dim_start, dim_stop = im.tuple_get(0, dim_range), im.tuple_get(1, dim_range) # bounds @@ -412,7 +410,7 @@ def _construct_itir_domain_arg( domain_args.append( itir.FunCall( fun=itir.SymRef(id="named_range"), - args=[itir.AxisLiteral(value=dim.tag, kind=dim.kind), lower, upper], + args=[itir.AxisLiteral(value=dim.tag), lower, upper], ) ) diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index dced958782..6cfebfd529 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -216,12 +216,6 @@ def skip_value( NamedFieldIndices: TypeAlias = Mapping[Tag, FieldIndex | SparsePositionEntry] -# Magic local dimension for the result of a `make_const_list`. -# A clean implementation will probably involve to tag the `make_const_list` -# with the neighborhood it is meant to be used with. -_CONST_DIM = common.ConstListDim - - @runtime_checkable class ItIterator(Protocol): """ @@ -570,7 +564,7 @@ def execute_shift( for i, p in reversed(list(enumerate(new_entry))): # first shift applies to the last sparse dimensions of that axis type if p is None: - if tag == _CONST_DIM.tag: + if tag == common.ConstList.tag: new_entry[i] = 0 else: # NOTE: the sparse tag is the local dimension's; the table over it may be @@ -1014,7 +1008,7 @@ def field_setitem(self, named_indices: NamedFieldIndices, value: Any): ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ - self._translate_named_indices({**named_indices, _CONST_DIM.tag: 0}) + self._translate_named_indices({**named_indices, common.ConstList.tag: 0}) ] = value.value else: self._ndarrayfield[self._translate_named_indices(named_indices)] = value @@ -1451,7 +1445,7 @@ def __gt_type__(self) -> ts.ListType: assert isinstance(element_type, ts.DataType) return ts.ListType( element_type=element_type, - offset_type=_CONST_DIM, + offset_type=common.ConstList, ) @@ -1548,7 +1542,7 @@ class SparseListIterator: offsets: Sequence[OffsetPart] = dataclasses.field(default_factory=list, kw_only=True) def deref(self) -> Any: - if self.list_offset == _CONST_DIM.tag: + if self.list_offset == common.ConstList.tag: return _ConstList( value=self.it.shift(*self.offsets, SparseTag(self.list_offset), 0).deref() ) @@ -1806,9 +1800,9 @@ def _fieldspec_list_to_value( ) -> tuple[common.Domain, ts.TypeSpec]: """Translate the list element type into the domain.""" if isinstance(type_, ts.ListType): - if type_.offset_type == _CONST_DIM: + if type_.offset_type is common.ConstList: return domain.insert( - len(domain), common.named_range((_CONST_DIM, 1)) + len(domain), common.named_range((common.ConstList, 1)) ), type_.element_type else: offset_provider = embedded_context.get_offset_provider() diff --git a/src/gt4py/next/iterator/ir.py b/src/gt4py/next/iterator/ir.py index f024ef6168..05561c2398 100644 --- a/src/gt4py/next/iterator/ir.py +++ b/src/gt4py/next/iterator/ir.py @@ -90,10 +90,19 @@ class OffsetLiteral(Expr): class AxisLiteral(Expr): - # TODO(havogt): Refactor to use declare Axis/Dimension at the Program level. - # Now every use of the literal has to provide the kind, where usually we only care of the name. + #: The dimension's tag, its qualified Python name (ADR 0028). value: str - kind: common.DimensionKind = common.DimensionKind.HORIZONTAL + + @property + def dim(self) -> common.Dimension: + """The dimension the literal names, resolved from its tag.""" + return common.resolve(self.value) + + @property + def kind(self) -> 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). + return self.dim.kind class CartesianOffset(Expr): diff --git a/src/gt4py/next/iterator/ir_utils/ir_makers.py b/src/gt4py/next/iterator/ir_utils/ir_makers.py index d6458f7a88..053e7b78c8 100644 --- a/src/gt4py/next/iterator/ir_utils/ir_makers.py +++ b/src/gt4py/next/iterator/ir_utils/ir_makers.py @@ -583,7 +583,7 @@ def _impl(*its: itir.Expr) -> itir.FunCall: def axis_literal(dim: common.Dimension) -> itir.AxisLiteral: - return itir.AxisLiteral(value=dim.tag, kind=dim.kind) + return itir.AxisLiteral(value=dim.tag) def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: @@ -640,7 +640,7 @@ def index(dim: common.Dimension) -> itir.FunCall: Returns: A function that constructs a Field of indices in the given dimension. """ - return call("index")(itir.AxisLiteral(value=dim.tag, kind=dim.kind)) + return call("index")(itir.AxisLiteral(value=dim.tag)) def map_list(op): diff --git a/src/gt4py/next/iterator/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index 07033ff1db..1fd5f557bc 100644 --- a/src/gt4py/next/iterator/pretty_parser.py +++ b/src/gt4py/next/iterator/pretty_parser.py @@ -41,7 +41,7 @@ // suffix terminates it. TAG: /[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*(?:\[[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*\])?/ OFFSET_LITERAL: ( INT_LITERAL | TAG ) "ₒ" - AXIS_LITERAL: TAG ("ᵥ" | "ₕ") + AXIS_LITERAL: TAG ("ᵥ" | "ₕ" | "ₗ") INFINITY_LITERAL: "∞" | "-∞" _literal: INT_LITERAL | FLOAT_LITERAL | OFFSET_LITERAL | AXIS_LITERAL | INFINITY_LITERAL ID_NAME: CNAME @@ -177,9 +177,8 @@ def INFINITY_LITERAL(self, value: lark_lexer.Token) -> ir.InfinityLiteral: return ir.InfinityLiteral.POSITIVE def AXIS_LITERAL(self, value: lark_lexer.Token) -> ir.AxisLiteral: - name = value.value[:-1] - kind = ir.DimensionKind.HORIZONTAL if value.value[-1] == "ₕ" else ir.DimensionKind.VERTICAL - return ir.AxisLiteral(value=name, kind=kind) + # NOTE: the kind suffix is only for the reader; the kind is the dimension's own. + return ir.AxisLiteral(value=value.value[:-1]) def lam(self, *args: ir.Node) -> ir.Lambda: *params, expr = args diff --git a/src/gt4py/next/iterator/pretty_printer.py b/src/gt4py/next/iterator/pretty_printer.py index 5fbba8920e..5321c01089 100644 --- a/src/gt4py/next/iterator/pretty_printer.py +++ b/src/gt4py/next/iterator/pretty_printer.py @@ -19,6 +19,7 @@ from typing import Final from gt4py.eve import NodeTranslator +from gt4py.next import common from gt4py.next.iterator import ir from gt4py.next.type_system import type_specifications as ts, type_translation @@ -134,6 +135,13 @@ def implied_literal_type(value: str) -> ts.ScalarType: DEFAULT_WIDTH: Final = 100 +_AXIS_KIND_SUFFIX: Final = { + common.DimensionKind.HORIZONTAL: "ₕ", + common.DimensionKind.VERTICAL: "ᵥ", + common.DimensionKind.LOCAL: "ₗ", +} + + class PrettyPrinter(NodeTranslator): def __init__( self, @@ -225,11 +233,11 @@ def visit_CartesianOffset(self, node: ir.CartesianOffset, *, prec: int) -> list[ return [f"{domain}→{codomain}"] def visit_AxisLiteral(self, node: ir.AxisLiteral, *, prec: int) -> list[str]: - kind = "" - if node.kind == ir.DimensionKind.HORIZONTAL: + try: + kind = _AXIS_KIND_SUFFIX[node.kind] + except ValueError: + # a tag that names no importable dimension, e.g. in IR built by hand for debugging kind = "ₕ" - elif node.kind == ir.DimensionKind.VERTICAL: - kind = "ᵥ" return [str(node.value) + kind] def visit_SymRef(self, node: ir.SymRef, *, prec: int) -> list[str]: diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index ec8db888b1..fb4eeddb37 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -141,7 +141,7 @@ def make_node(o): if isinstance(o, Node): return o if isinstance(o, common.DimensionMeta): - return AxisLiteral(value=o.tag, kind=o.kind) + return AxisLiteral(value=o.tag) if isinstance(o, common.Infinity): if o is common.Infinity.POSITIVE: return itir.InfinityLiteral.POSITIVE diff --git a/src/gt4py/next/iterator/transforms/remove_broadcast.py b/src/gt4py/next/iterator/transforms/remove_broadcast.py index db6445cc51..1a92361216 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -34,9 +34,7 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): >>> domain = im.domain(common.GridType.CARTESIAN, {IDim: (0, 10), JDim: (0, 10)}) >>> expr = im.call("broadcast")( ... im.ref("inp"), - ... im.make_tuple( - ... *(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (IDim, JDim)) - ... ), + ... im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (IDim, JDim))), ... ) >>> expr.annex.domain = domain_utils.SymbolicDomain.from_expr(domain) >>> transformed = RemoveBroadcast.apply(expr) 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 3b44ecc4c2..8a8e3fce53 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 @@ -58,13 +58,6 @@ ) -# Magic local dimension used for list of values with length known at compile-time. -# NOTE: the canonical class from `common`, not a local declaration: under nominal identity a -# second declaration would be a *different* dimension and the `== _CONST_DIM` checks below -# would stop matching `ListType`s built by embedded execution. -_CONST_DIM: Final = gtx_common.ConstListDim - - @dataclasses.dataclass(frozen=True) class ValueExpr: """ @@ -595,7 +588,7 @@ def _construct_tasklet_result( return ValueExpr( dc_node=temp_node, gt_dtype=( - ts.ListType(element_type=data_type, offset_type=_CONST_DIM) + ts.ListType(element_type=data_type, offset_type=gtx_common.ConstList) if use_array else data_type ), @@ -1144,7 +1137,8 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: gt_field=ts.FieldType( dims=[conn_type.domain[0]], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1320,7 +1314,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: assert isinstance(input_arg.gt_dtype, ts.ListType) assert input_arg.gt_dtype.offset_type is not None offset_type = input_arg.gt_dtype.offset_type - if offset_type == _CONST_DIM: + 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( @@ -1363,7 +1357,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: raise ValueError(f"More than one local dimension in map expression {node}.") input_size = input_desc.shape[0] if input_size == 1: - assert input_arg.gt_dtype.offset_type == _CONST_DIM + assert input_arg.gt_dtype.offset_type is gtx_common.ConstList input_memlets[conn] = dace.Memlet(data=input_node.data, subset="0") elif input_size == local_size: input_memlets[conn] = dace.Memlet(data=input_node.data, subset=map_index) @@ -1396,7 +1390,8 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: gt_field=ts.FieldType( dims=[conn_type.domain[0]], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1726,7 +1721,8 @@ def _make_unstructured_shift( gt_field=ts.FieldType( dims=[conn_type.source_dim], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1931,7 +1927,7 @@ def _visit_Lambda_impl( and node.expr.type.offset_type is not None and isinstance(result, (MemletExpr, ValueExpr)) and isinstance(result.gt_dtype, ts.ListType) - and result.gt_dtype.offset_type == _CONST_DIM + and result.gt_dtype.offset_type is gtx_common.ConstList ): result = self._broadcast_const_list(result, node.expr.type) 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 18358ac418..6d423a52eb 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 @@ -980,7 +980,7 @@ def foo(inp: gtx.Field[[TDim], float64]): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( im.ref("inp"), - im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (TDim, UDim))), ) @@ -994,7 +994,7 @@ def foo(): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( 1, - im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (TDim, UDim))), ) 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 be0048c285..842e84fc4e 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 @@ -10,6 +10,7 @@ import pytest import gt4py.next as gtx +from gt4py.next import common from gt4py.next.embedded import context as embedded_context from gt4py.next.iterator import embedded, runtime from gt4py.next.iterator.builtins import ( @@ -69,7 +70,7 @@ def testee(): ref = np.asarray([[42.0], [42.0]]) assert result.domain.dims[0] == E - assert result.domain.dims[1] == embedded._CONST_DIM # this is implementation detail + assert result.domain.dims[1] == common.ConstList # this is implementation detail assert result.shape[1] == 1 # this is implementation detail np.testing.assert_array_equal(result.asnumpy(), ref) @@ -152,6 +153,6 @@ def testee(): ref = np.asarray([[43.0], [43.0]]) assert result.domain.dims[0] == E - assert result.domain.dims[1] == embedded._CONST_DIM # this is implementation detail + assert result.domain.dims[1] == common.ConstList # this is implementation detail assert result.shape[1] == 1 # this is implementation detail np.testing.assert_array_equal(result.asnumpy(), ref) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py index 83a16a980c..52b358023d 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py @@ -190,7 +190,7 @@ def test_named_range_unbounded(): expected = ir.FunCall( fun=ir.SymRef(id="named_range"), args=[ - ir.AxisLiteral(value="KDim", kind=ir.DimensionKind.VERTICAL), + ir.AxisLiteral(value="KDim"), ir.InfinityLiteral.NEGATIVE, ir.InfinityLiteral.POSITIVE, ], @@ -257,7 +257,7 @@ def test_named_range_vertical(): expected = ir.FunCall( fun=ir.SymRef(id="named_range"), args=[ - ir.AxisLiteral(value="IDim", kind=ir.DimensionKind.VERTICAL), + ir.AxisLiteral(value="IDim"), ir.SymRef(id="x"), ir.SymRef(id="y"), ], 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 1144b727b2..0b41235e56 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 @@ -8,6 +8,8 @@ import pytest +import gt4py.next as gtx + from gt4py.next.iterator import builtins, ir, pretty_printer from gt4py.next.iterator.ir_utils import ir_makers as im from gt4py.next.iterator.pretty_printer import PrettyPrinter, pformat @@ -259,18 +261,23 @@ def test_make_tuple(): assert actual == expected -def test_axis_literal_horizontal(): - testee = ir.AxisLiteral(value="I", kind=ir.DimensionKind.HORIZONTAL) - expected = "Iₕ" - actual = pformat(testee) - assert actual == expected +class IDim(gtx.DimensionIndex): ... -def test_axis_literal_vertical(): - testee = ir.AxisLiteral(value="I", kind=ir.DimensionKind.VERTICAL) - expected = "Iᵥ" - actual = pformat(testee) - assert actual == expected +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class LocalDim(gtx.LocalDimensionIndex): ... + + +@pytest.mark.parametrize("dim, suffix", [(IDim, "ₕ"), (KDim, "ᵥ"), (LocalDim, "ₗ")]) +def test_axis_literal(dim, suffix): + # the suffix is the resolved dimension's kind + assert pformat(ir.AxisLiteral(value=dim.tag)) == f"{dim.tag}{suffix}" + + +def test_axis_literal_of_unresolvable_tag(): + assert pformat(ir.AxisLiteral(value="I")) == "Iₕ" def test_named_range_horizontal(): diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py index 81396d8c59..6584028a93 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py @@ -72,19 +72,8 @@ im.tuple_get(im.literal("42", builtins.INTEGER_INDEX_BUILTIN), "x"), id="tuple_get" ), pytest.param(im.make_tuple("x", "y"), id="make_tuple"), - pytest.param( - ir.AxisLiteral(value="I", kind=ir.DimensionKind.HORIZONTAL), id="axis_literal_horizontal" - ), - pytest.param( - ir.AxisLiteral(value="I", kind=ir.DimensionKind.VERTICAL), id="axis_literal_vertical" - ), - pytest.param( - im.named_range(ir.AxisLiteral(value="IDim"), "x", "y"), id="named_range_horizontal" - ), - pytest.param( - im.named_range(ir.AxisLiteral(value="IDim", kind=ir.DimensionKind.VERTICAL), "x", "y"), - id="named_range_vertical", - ), + pytest.param(ir.AxisLiteral(value="I"), id="axis_literal"), + pytest.param(im.named_range(ir.AxisLiteral(value="IDim"), "x", "y"), id="named_range"), pytest.param(im.call("cartesian_domain")("x"), id="cartesian_domain"), pytest.param(im.call("unstructured_domain")("x"), id="unstructured_domain"), pytest.param(im.if_("x", "y", "z"), id="if_short"), @@ -161,7 +150,7 @@ pytest.param(ir.InfinityLiteral.NEGATIVE, id="infinity_negative"), pytest.param( im.named_range( - ir.AxisLiteral(value="KDim", kind=ir.DimensionKind.VERTICAL), + ir.AxisLiteral(value="KDim"), ir.InfinityLiteral.NEGATIVE, im.literal("5", "int32"), ), diff --git a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py index 94f0610c06..4f580b110c 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py @@ -101,9 +101,7 @@ def expression_test_cases(): bool_type, ), ( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), it_ts.NamedRangeType(dim=Vertex), ), ( @@ -112,9 +110,7 @@ def expression_test_cases(): ), ( im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ) + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1) ), ts.DomainType(dims=[Vertex]), ), @@ -443,10 +439,8 @@ def test_cartesian_fencil_definition(): def test_unstructured_fencil_definition(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), - im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( @@ -510,10 +504,8 @@ def test_function_definition(): def test_fencil_with_nb_field_input(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), - im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( From e00e4a927d7675d4149ed07aeb5e5f7f92e5926a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 06:28:30 +0200 Subject: [PATCH 15/17] fix[next]: printing IR does not import modules The pretty printer takes an axis literal's kind from its inferred type, or from a dimension that is already loaded (common.resolve_loaded), and never imports. gtfn's domain canonicalization resolves through dim_from_axis_literal; ADR 0028 states what resolve memoizes. --- .../next/0028-Dimensions_As_Nominal_Types.md | 6 ++++-- src/gt4py/next/common.py | 20 +++++++++++++++++++ src/gt4py/next/iterator/pretty_parser.py | 2 +- src/gt4py/next/iterator/pretty_printer.py | 15 ++++++++------ .../codegens/gtfn/itir_to_gtfn_ir.py | 9 +++++++-- .../iterator_tests/test_pretty_parser.py | 3 ++- .../iterator_tests/test_pretty_printer.py | 8 ++++++++ 7 files changed, 51 insertions(+), 12 deletions(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 295ceebe59..28cee31a28 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -77,8 +77,10 @@ disappears. 1. **Reconstruction from the IR is an import.** `common.resolve(tag)` imports the module and walks the qualname; nested declarations resolve naturally. The IR - references a Python type exactly the way `pickle` references a class. It is - memoized, because type inference calls it once per `AxisLiteral`. An + references a Python type exactly the way `pickle` references a class. Where the + module path ends is memoized, because type inference calls it once per + `AxisLiteral`; the attribute walk is repeated, so a redefined declaration is + found. An `AxisLiteral` stores only the tag: its `kind` is the resolved dimension's, so the two cannot disagree. diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index ef43d5593c..c88ea5d4c1 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -411,6 +411,26 @@ def _import_qualified_name_or_none(tag: Tag) -> Any: return obj +def resolve_loaded(tag: Tag) -> Optional[Dimension]: + """ + Return the dimension a tag names if its module is already loaded, else `None`. + + Like `resolve`, but never imports: for code that must not have import side effects, such as + printing IR. + """ + if (match := _STAGGERED_TAG_RE.match(tag)) is not None: + owner, base = resolve_loaded(match["owner"]), resolve_loaded(match["base"]) + return owner[base] if owner is not None and base is not None else None # type: ignore[index] # parametrized dimension + parts = tag.split(".") + for split in range(len(parts) - 1, 0, -1): + if (obj := sys.modules.get(".".join(parts[:split]))) is None: + continue + for attr in parts[split:]: + obj = getattr(obj, attr, None) + return obj if isinstance(obj, DimensionMeta) else None + return None + + def _import_qualified_name(tag: Tag) -> Any: """Import the object a dotted qualified name refers to; see `resolve`.""" module_name, attrs = _split_qualified_name(tag) diff --git a/src/gt4py/next/iterator/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index 1fd5f557bc..0324499d26 100644 --- a/src/gt4py/next/iterator/pretty_parser.py +++ b/src/gt4py/next/iterator/pretty_parser.py @@ -20,7 +20,7 @@ from gt4py.next.type_system import type_specifications as ts -GRAMMAR = """ +GRAMMAR = r""" start: fencil_definition | function_definition | declaration diff --git a/src/gt4py/next/iterator/pretty_printer.py b/src/gt4py/next/iterator/pretty_printer.py index 5321c01089..a7376e0d22 100644 --- a/src/gt4py/next/iterator/pretty_printer.py +++ b/src/gt4py/next/iterator/pretty_printer.py @@ -16,7 +16,7 @@ import types as _types from collections.abc import Iterator, Mapping, Sequence -from typing import Final +from typing import Final, Optional from gt4py.eve import NodeTranslator from gt4py.next import common @@ -233,11 +233,14 @@ def visit_CartesianOffset(self, node: ir.CartesianOffset, *, prec: int) -> list[ return [f"{domain}→{codomain}"] def visit_AxisLiteral(self, node: ir.AxisLiteral, *, prec: int) -> list[str]: - try: - kind = _AXIS_KIND_SUFFIX[node.kind] - except ValueError: - # a tag that names no importable dimension, e.g. in IR built by hand for debugging - kind = "ₕ" + # NOTE: printing must not import modules (`str()` of any node prints it), so the kind is + # taken from the inferred type, or from an already loaded dimension. A tag naming neither, + # e.g. in IR built by hand, prints as horizontal; the parser ignores the suffix anyway. + if isinstance(node.type, ts.DimensionType): + 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 "ₕ" return [str(node.value) + kind] def visit_SymRef(self, node: ir.SymRef, *, prec: int) -> list[str]: 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 74684d420b..38fe8fa1c6 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 @@ -254,11 +254,16 @@ def visit_FunCall(self, node: itir.FunCall) -> itir.FunCall: assert isinstance(node.args[0], itir.FunCall) first_axis_literal = node.args[0].args[0] assert isinstance(first_axis_literal, itir.AxisLiteral) - if first_axis_literal.kind == itir.DimensionKind.VERTICAL: + if ir_utils_misc.dim_from_axis_literal(first_axis_literal).kind == ( + itir.DimensionKind.VERTICAL + ): assert len(node.args) == 2 assert isinstance(node.args[1], itir.FunCall) assert isinstance(node.args[1].args[0], itir.AxisLiteral) - assert node.args[1].args[0].kind == itir.DimensionKind.HORIZONTAL + assert ( + ir_utils_misc.dim_from_axis_literal(node.args[1].args[0]).kind + == itir.DimensionKind.HORIZONTAL + ) return itir.FunCall(fun=node.fun, args=[node.args[1], node.args[0]]) return node diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py index 52b358023d..7165a54ec2 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py @@ -252,7 +252,8 @@ def test_named_range_horizontal(): assert actual == expected -def test_named_range_vertical(): +def test_named_range_kind_suffix_is_ignored(): + # the kind is the dimension's own; the suffix only helps the reader testee = "IDimᵥ: [x, y[" expected = ir.FunCall( fun=ir.SymRef(id="named_range"), 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 0b41235e56..c02336e6e0 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 @@ -277,7 +277,15 @@ def test_axis_literal(dim, suffix): def test_axis_literal_of_unresolvable_tag(): + # printing does not import modules: a tag naming no loaded dimension prints as horizontal, + # so text -> IR -> text is not the identity for such tags (the parser ignores the suffix) assert pformat(ir.AxisLiteral(value="I")) == "Iₕ" + assert pformat(ir.AxisLiteral(value="this.I")) == "this.Iₕ" + + +def test_axis_literal_kind_from_type(): + typed = ir.AxisLiteral(value="not.loaded.KDim", type=ts.DimensionType(dim=KDim)) + assert pformat(typed) == "not.loaded.KDimᵥ" def test_named_range_horizontal(): From 2a468c69a2908a6a68928a0dc632e2b5ed3fbf6e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 07:19:43 +0200 Subject: [PATCH 16/17] test[next]: register doctest dimensions where their tags point --- .../transforms/replace_get_domain_range_with_constants.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py index 3714ea56ad..a9b3a35386 100644 --- a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py +++ b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py @@ -56,6 +56,9 @@ class ReplaceGetDomainRangeWithConstants(PreserveLocationVisitor, NodeTranslator >>> from gt4py import next as gtx >>> class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... >>> class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + >>> import sys # register the dimensions where their tags point, as a module would + >>> sys.modules[__name__].KDim = KDim + >>> sys.modules[__name__].Vertex = Vertex >>> sizes = { ... "out": gtx.domain({Vertex: (0, 10), KDim: (0, 20)}), From 36cbcb7d2790decbf4dd70179267cfc91a2425fa Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 16:14:44 +0200 Subject: [PATCH 17/17] fix[next]: review fixes for ConstList and AxisLiteral - rename the references the earlier PRs added along with the class - ADR 0028's date, which this PR's edits had outrun - say why 'ListType.offset_type' can be 'None' while embedded uses 'ConstList' - test 'resolve_loaded', which is what keeps printing IR import-free --- .../ADRs/next/0028-Dimensions_As_Nominal_Types.md | 2 +- src/gt4py/next/common.py | 4 ++-- src/gt4py/next/type_system/type_specifications.py | 4 ++++ tests/next_tests/unit_tests/test_common.py | 9 +++++++++ .../next_tests/unit_tests/test_neighbor_connectivity.py | 4 ++-- 5 files changed, 18 insertions(+), 5 deletions(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 28cee31a28..35c2c2805b 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -7,7 +7,7 @@ tags: [] - **Status**: proposed - **Authors**: Enrique González Paredes (@egparedes) - **Created**: 2026-09-18 -- **Updated**: 2026-09-18 +- **Updated**: 2026-09-23 A concrete dimension becomes a **class**, and an index along it an **instance** of that class — the shape `enum.Enum` uses, where the class is the collection and diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index c88ea5d4c1..64c6e58a3c 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -2238,9 +2238,9 @@ def __init_subclass__( " ('class Local(LocalDimensionIndex): ...') or by adopting one" " ('Local: TypeAlias = SomeLocalDim')." ) - if local is ConstListDim: + if local is ConstList: raise TypeError( - f"'{name}' cannot adopt '{ConstListDim.__qualname__}': it is the local dimension" + 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) diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index f5e0ea770c..afe5b350da 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -119,6 +119,10 @@ class ListType(DataType): """ element_type: DataType + #: The local dimension the list runs along. `None` where type inference does not know it, + #: which is how it spells the result of `make_const_list`; embedded execution and the DaCe + #: lowering use `common.ConstList` for the same thing. + #: TODO(egparedes): use `common.ConstList` in type inference too, and drop `None`. offset_type: common.Dimension | None diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index d1b9f97560..6629711027 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -918,3 +918,12 @@ def test_a_dimension_cannot_be_staggered_twice(self): def test_resolve_rejects_a_bracketed_tag_of_another_owner(self): with pytest.raises(ValueError, match="not a parametrized dimension"): common.resolve(f"{KDim.tag}[{KDim.tag}]") + + +def test_resolve_loaded(): + # `resolve_loaded` never imports: it answers for what is loaded and gives up otherwise + assert common.resolve_loaded(IDim.tag) is IDim + assert common.resolve_loaded(common.Staggered[IDim].tag) is common.Staggered[IDim] + assert common.resolve_loaded("not_imported_anywhere.IDim") is None + assert common.resolve_loaded(f"{__name__}.does_not_exist") is None + assert common.resolve_loaded(f"{__name__}.test_resolve_loaded") is None # not a dimension diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 2aaa12c46e..7462d38d2b 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -68,7 +68,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 @@ -649,6 +649,6 @@ def test_the_const_list_dimension_cannot_be_adopted(): _declare( """ class C(NeighborConnectivity[Vertex, Edge]): - Local: typing.TypeAlias = ConstListDim + Local: typing.TypeAlias = ConstList """ )