diff --git a/docs/development/ADRs/next/0019-Connectivities.md b/docs/development/ADRs/next/0019-Connectivities.md index 21827941bf..0c80a6cf55 100644 --- a/docs/development/ADRs/next/0019-Connectivities.md +++ b/docs/development/ADRs/next/0019-Connectivities.md @@ -31,7 +31,7 @@ We update and introduce the following concepts **NeighborTable** is a _GatherConnectivity_ that is a 2D mapping of the N neighbors of a Location A to a Location B, backed by a buffer. -**ConnectivityType**, **NeighborConnectivityType** contains all information that is needed for compilation. +**ConnectivityType**, **NeighborTableType** contain all information that is needed for compilation. A `NeighborTableType` is the type of a table bound to a `NeighborConnectivity` declaration (ADR 0030). ### Full definitions @@ -48,7 +48,7 @@ Embedded execution of iterator (local) view supports only `NeighborTable`s. ### IR transformations and compiled backends -All transformations and code-generation should use `ConnectivityType`, not the `Connectivity` which contains the runtime mapping. +All transformations and code-generation should use `ConnectivityType` / `NeighborTableType`, not the `Connectivity` which contains the runtime mapping. Note, currently the `global_tmps` pass uses runtime information, therefore this is not strictly enforced. @@ -60,3 +60,7 @@ The only supported `Connectivity`s in compiled backends (currently) are `Neighbo - Removed the abstract `NeighborConnectivity` concept; `NeighborTable` is now the single neighbor-connectivity concept (there is no non-buffer-backed neighbor connectivity in use). - Added `GatherConnectivity` (a `Connectivity` whose `premap` rearranges data via a gather), which the embedded field-view `premap` dispatches on. It replaces the former `ConnectivityKind` flag and unifies the previous reshuffling/remapping `premap` implementations into a single advanced-index gather. + +### 2026-09-24 + +- `NeighborConnectivityType` is renamed `NeighborTableType` and typed by the `NeighborConnectivity` declaration its table is bound to; a `NeighborTable`'s own `__gt_type__()` is the structural `ConnectivityType` (ADR 0030). diff --git a/docs/development/ADRs/next/0030-Connectivities_As_Types.md b/docs/development/ADRs/next/0030-Connectivities_As_Types.md new file mode 100644 index 0000000000..07ea6992b4 --- /dev/null +++ b/docs/development/ADRs/next/0030-Connectivities_As_Types.md @@ -0,0 +1,225 @@ +--- +tags: [] +--- + +# Connectivities as Types + +- **Status**: proposed +- **Authors**: Enrique González Paredes (@egparedes) +- **Created**: 2026-09-21 +- **Updated**: 2026-10-02 + +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 a neighbor table bound to it has to satisfy. It holds no data. It builds on [ADR 0029](0029-Dimensions_As_Nominal_Types.md): the +connectivity, like a dimension, is identified by its type, and `V2E.Local` is an +ordinary dimension class with the tag `.V2E.Local`. + +## 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[Domain, Codomain]` is a PEP 695 generic whose subclasses + are declarations: for each `Domain` element, a list of `Codomain` neighbors. Its + metaclass, `ConnectivityMeta`, forbids instantiation. The two dimensions are + the class attributes `V2E.domain` and `V2E.codomain`. +- The local dimension is the nested class `Local`, a subclass of + `LocalDimensionIndex`. Declaring it is required, and `NeighborConnectivity` + sets `Local.owner` to the connectivity when the class is created. A local + 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, 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. A sharer must have + the owner's domain; its codomain is free. The local dimension records its + sharers (`Local.sharers`) as it records its owner. +- `max_neighbors` and `min_neighbors` are optional class keywords, not type + parameters: Python has no integer type parameters, and nothing static needs + the count. A declared count is a constraint on the bound table; an undeclared + one is taken from the table. `min_neighbors < max_neighbors` means that the + table must use skip values. +- `common.check_neighbor_table(V2E, table)` checks a table, or just its type + (which is all an ahead-of-time compilation has), against the declaration, and + returns the table's `NeighborTableType` (below): the domain is + `(Domain, V2E.Local)`, the codomain is `Codomain`, the dtype is integral, and + the neighbor counts and skip values agree. Skip values are checked on the + table's type: a table with a `skip_value` counts as having skip values whether + or not an entry uses it. + +`Domain` and `Codomain` name the two index spaces the declaration maps between. +A bound table is a field over `(Domain, Local)` with values in `Codomain`: the +table's domain is the declaration's domain extended by the local axis, which is +the same use of the word as `Connectivity.domain` and +`CartesianConnectivity.domain_dim`. "Origin" would have been the other natural +name for the first dimension, but gt4py already uses it for the start of a +buffer (`__gt_origin__`). + +### The type of a bound table + +Transformations and code generation see types, never tables (ADR 0019). The type +of a table bound to a declaration is a `common.NeighborTableType`: +`connectivity` (the declaration), `dtype`, `skip_value` and `max_neighbors`. Its +`domain` and `codomain` are derived from the declaration, +`(connectivity.domain, local_dimension_of(connectivity))` and +`connectivity.codomain`, so they cannot disagree with it. The mapping from +offset-provider keys to these records is `common.TableTypes`, and it can be +given instead of the tables for ahead-of-time compilation. + +A table cannot tell which declaration it is bound to: the table of a sharer +(`C2CE`) has the same domain as its owner's (`C2E`), with another codomain. So a +`NeighborTableType` is built where a table is bound, from its offset-provider key: +`check_neighbor_table(C2CE, table)`, or `offset_provider_to_type`, which finds +the declaration whose `offset_tag` is the key among the owner and the sharers of +the table's local dimension. `NeighborTable.__gt_type__()` returns only what the +table knows, the structural `common.ConnectivityType` (domain, codomain, dtype, +skip value). + +A table bound under a key that no declaration answers to -- hand-written IR +names its offsets by plain strings -- has no declaration. Its +`NeighborTableType` then has the table's structural `ConnectivityType` as its +`connectivity`, and `domain` and `codomain` are read from that. This keeps the +IR level, which does not know declarations, working unchanged. + +A `NeighborTableType` is fingerprinted through its fields, so the declaration +takes part in the fingerprint of everything compiled for it: the owner's and a +sharer's tables, identical as tables, produce different artifact keys. + +### `NeighborConnectivity` is not a `Connectivity` + +`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 tells local dimensions apart at runtime, and constructors whose parameter +must be a primary dimension (`NeighborConnectivity[Domain, Codomain]`) check it; +`Staggered[D]` rejects a local dimension statically too, through its bound on a +declared Cartesian axis (ADR 0029). + +### Localness is the class, and `DimensionKind.LOCAL` is removed + +A dimension is local if and only if it subclasses `LocalDimensionIndex` +(`common.is_local_dimension(dim)`), so `DimensionKind.LOCAL` is removed and +`DimensionKind` is `HORIZONTAL | VERTICAL`. A local dimension's `kind` is `None`, +and declaring one with `kind=` is a `TypeError`. `None` rather than `HORIZONTAL` +keeps every `kind == HORIZONTAL` / `kind != VERTICAL` comparison in the backends +meaning what it meant: with `HORIZONTAL`, a sparse field would silently count its +local axis as horizontal. Two consequences: + +- `None` does not order against the enum, so `order_dimensions` sorts by an + explicit rank — horizontal, then local, then vertical — the order `kind` used to + encode. It must not move, or the memory layout of sparse fields changes with it. +- Displays derive the label from the class: `str(V2E.Local)` is still + `Local[local]`, the IR pretty printer still marks a local axis with `ₗ`, and DaCe + map variables of a local dimension keep their `_gtx_localdim` suffix. + +What remains of `kind` is the layout sort key and the scan axis. + +### `Local` is not annotated anywhere + +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 + +A declaration is typed like the `FieldOffset` it replaces: `V2E.__gt_type__()` +is a `ts.ShiftType`, which takes a field over the codomain to one over the +domain, `Shift[: Edge -> (Vertex, V2E.Local)]`. `V2E[i]` has the domain +`(Vertex,)`, and so does a Cartesian shift `KDim + 1`, over `KDim` and without a +tag. The tag is the connectivity's `offset_tag`: + +- **the local dimension's tag**, `V2E.Local.tag`, for the connectivity that + declares it. This is the single string that shifts, neighbor reductions and + 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 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. + +`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 + +- An unstructured connectivity is spelled once. The provider key, the offset tag + and the local dimension are all derived from the declaration. +- A table bound to a connectivity can be checked against its declaration. +- `V2E.Local` in DSL code is resolved from the shift type, because the type of + `V2E` is not the class. +- Code generation sees which declaration a table is bound to, not only its + shape; a table without a declaration is typed by its structure. +- A declaration is fingerprinted by its name *and* its declared dimensions and + counts, so redefining it under the same name (e.g. re-running a notebook + cell) does not reuse artifacts compiled for the old declaration. +- `FieldOffset` remains during migration; a `FieldOffset` and a + `NeighborConnectivity` sharing a local dimension are interchangeable. + +## Alternatives considered + +- **The local dimension generated by the metaclass**, e.g. `V2E.Local` created + from `V2E`'s name. Type checkers cannot see a generated class, so it could not + be used in `Field[Dims[Vertex, V2E.Local], ...]`. +- **Neighbor counts as type parameters.** Python has no integer type parameters, + and a `Literal[6]` argument would add a type parameter nothing statically uses. +- **`NeighborConnectivity` as a `Connectivity` subclass.** Mixes the + declaration with the data protocol; see above. diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index 4f1661d745..35c94b5b78 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) - [0029 - Dimensions as Nominal Types](0029-Dimensions_As_Nominal_Types.md) +- [0030 - Connectivities as Types](0030-Connectivities_As_Types.md) ### Frontend and Parsing #frontend diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 63d5c2dda3..1695c0a63f 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -227,7 +227,7 @@ Another way to look at it is that transform uses the edge-to-cell connectivity t You can use the field offset `E2C` below to transform a field over cells to a field over edges using the edge-to-cell connectivities: ```{code-cell} ipython3 -class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2CDim(gtx.LocalDimensionIndex): ... E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim)) ``` @@ -379,7 +379,7 @@ print("where nested tuple return: {}".format(((result_1.asnumpy(), result_2.asnu As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to define the local dimension, the field offset, and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: ```{code-cell} ipython3 -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(gtx.LocalDimensionIndex): ... C2E = gtx.FieldOffset(C2EDim.tag, source=EdgeDim, target=(CellDim, C2EDim)) C2E_offset_provider = gtx.as_connectivity([CellDim, C2EDim], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index 82c0ef6a57..272fbdcc57 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -11,7 +11,14 @@ import gt4py.next as gtx from gt4py.next.iterator.embedded import MutableLocatedField from gt4py.next import neighbor_sum, where, Dims -from gt4py.next import CartesianAxisIndex, Dimension, DimensionIndex, DimensionKind, FieldOffset +from gt4py.next import ( + CartesianAxisIndex, + Dimension, + DimensionIndex, + LocalDimensionIndex, + DimensionKind, + FieldOffset, +) from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( run_gtfn as gtfn_cpu, @@ -389,31 +396,31 @@ class E(DimensionIndex): ... class K(CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... C2E = FieldOffset(C2EDim.tag, source=E, target=(C, C2EDim)) -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... V2E = FieldOffset(V2EDim.tag, source=E, target=(V, V2EDim)) -class E2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2VDim(LocalDimensionIndex): ... E2V = FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) -class E2CDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2CDim(LocalDimensionIndex): ... E2C = FieldOffset(E2CDim.tag, source=C, target=(E, E2CDim)) -class E2C2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2VDim(LocalDimensionIndex): ... E2C2V = FieldOffset(E2C2VDim.tag, source=V, target=(E, E2C2VDim)) diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index 6a44a6eef7..a336ed664a 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -273,7 +273,7 @@ "metadata": {}, "outputs": [], "source": [ - "class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...\n", + "class E2CDim(gtx.LocalDimensionIndex): ...\n", "\n", "\n", "E2C = gtx.FieldOffset(E2CDim.tag, source=Cell, target=(Edge, E2CDim))" diff --git a/noxfile.py b/noxfile.py index c8091de09e..31327f076d 100755 --- a/noxfile.py +++ b/noxfile.py @@ -349,6 +349,9 @@ def test_typing_exports(session: nox.Session) -> None: "typing_tests", *session.posargs, ) + # A second checker, on code that must type-check for a downstream user: mypy and pyright + # disagree about what counts as a type, which the mypy-only cases above cannot catch. + session.run("pyright", "--project", "typing_tests", "typing_tests/pyright_probes.py") # -- DaCe codegen determinism check -- diff --git a/pyproject.toml b/pyproject.toml index e131ebda0e..ba7f6cd74d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -65,6 +65,7 @@ typing = [ typing_exports = [ # to test typing with gt4py in downstream code {include-group = "typing"}, + 'pyright>=1.1.400', # the second checker: it disagrees with mypy about what counts as a type 'types-six>=1.17.0.20251009', # can not let mypy auto-install types as that leads to unexpected stderr output (which means test failure) 'pytest-mypy-plugins>=4.0.0', # pytest plugin for running mypy on code snippets "xarray>=2024.1.0" # one of the regression tests requires xarray @@ -252,8 +253,6 @@ markers = [ 'uses_ir_if_stmts', 'uses_lift: tests that require backend support for lift builtin function', 'uses_negative_modulo: tests that require backend support for modulo on negative numbers', - 'uses_offset_tag_differing_from_local_dim: tests that shift with an offset tag that differs from the name of its local dimension', - 'uses_offset_tag_differing_from_local_dim_in_reduction: tests that reduce over an offset whose tag differs from the name of its local dimension', 'uses_origin: tests that require backend support for domain origin', 'uses_reduce_with_lambda: tests that use lambdas as reduce functions', 'uses_reduction_with_only_sparse_fields: tests that require backend support for with sparse fields', diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index b8a7bf5143..253e574f95 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -33,12 +33,15 @@ Domain, Field, GridType, + LocalDimensionIndex, + NeighborConnectivity, Staggered, UnitRange, as_non_staggered, domain, flip_staggered, is_staggered, + local_dimension_of, resolve, unit_range, ) @@ -124,6 +127,8 @@ "AnyCartesianAxisIndex", "CartesianAxisIndex", "DimensionKind", + "LocalDimensionIndex", + "NeighborConnectivity", "Staggered", "resolve", "Dims", @@ -136,6 +141,7 @@ "unit_range", "UnitRange", "is_staggered", + "local_dimension_of", "flip_staggered", "as_non_staggered", # from constructors diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index a5486c024f..471b4a7abb 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 @@ -39,6 +40,8 @@ TypeVarTuple, Unpack, cast, + get_args, + get_origin, overload, ) @@ -121,15 +124,36 @@ def from_codegen_name(name: str) -> Tag: @enum.unique class DimensionKind(StrEnum): + """ + The role of a non-local dimension: a field's layout sort key, and the scan axis. + + There is no `LOCAL` member: whether a dimension is local is a fact about its class + (`is_local_dimension`), and a local dimension's `kind` is `None` (ADR 0030). + """ + HORIZONTAL = "horizontal" VERTICAL = "vertical" - LOCAL = "local" def __str__(self) -> str: return self.value -_DIM_KIND_ORDER = {DimensionKind.HORIZONTAL: 0, DimensionKind.LOCAL: 1, DimensionKind.VERTICAL: 2} +def is_local_dimension(dim: Any) -> bool: + """Return whether `dim` is a local dimension, i.e. a subclass of `LocalDimensionIndex`.""" + return isinstance(dim, DimensionMeta) and issubclass(dim, LocalDimensionIndex) + + +def _dimension_rank(dim: Dimension) -> int: + # NOTE: an explicit rank rather than a sort on `kind`: a local dimension's `kind` is `None`, + # which does not order against the enum. Horizontal, then local, then vertical -- the order + # `kind` used to encode; changing it would change the memory layout of sparse fields. + if is_local_dimension(dim): + return 1 + return 2 if dim.kind is DimensionKind.VERTICAL else 0 + + +def _kind_label(dim: DimensionMeta) -> str: + return "local" if is_local_dimension(dim) else str(dim.kind) class DimensionMeta(type): @@ -141,7 +165,8 @@ class DimensionMeta(type): metaclass, so this is the only place they can live. """ - kind: DimensionKind + #: `None` for a local dimension, whose localness is its class (see `is_local_dimension`). + kind: Optional[DimensionKind] # NOTE: mandatory, not redundant. Python sets `__hash__ = None` on any class body that # defines `__eq__` without it -- metaclasses included -- and `__eq__` below stays for the @@ -177,12 +202,12 @@ def value(cls) -> NoReturn: ) def __repr__(cls) -> str: - return f"{cls.tag}[{cls.kind}]" + return f"{cls.tag}[{_kind_label(cls)}]" def __str__(cls) -> str: # NOTE: the unqualified name, so diagnostics stay readable. `tag` is identity, not a # display name; `repr` carries the module and disambiguates when it matters. - return f"{cls.__qualname__}[{cls.kind}]" + return f"{cls.__qualname__}[{_kind_label(cls)}]" # NOTE: the self-type restricts index arithmetic to a Cartesian axis for the type checkers: # both bind it correctly at every call site (`C + 1` is an error for a mesh location `C`), @@ -276,7 +301,7 @@ class DimensionIndex(metaclass=DimensionMeta): False """ - kind: ClassVar[DimensionKind] = DimensionKind.HORIZONTAL + kind: ClassVar[Optional[DimensionKind]] = DimensionKind.HORIZONTAL __slots__ = ("value",) @@ -1107,7 +1132,9 @@ def asnumpy(self) -> np.ndarray: ... def as_scalar(self) -> core_defs.ScalarT: ... @abc.abstractmethod - def premap(self, index_field: Connectivity | fbuiltins.FieldOffset) -> Field: ... + def premap( + self, index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity] + ) -> Field: ... @abc.abstractmethod def restrict(self, item: AnyIndexSpec) -> Self: ... @@ -1115,8 +1142,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 @@ -1298,17 +1325,44 @@ 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 +class NeighborTableType: + """ + The type of a neighbor table bound to a connectivity: what transformations and code generation + see instead of the table (ADR 0019). + + `connectivity` is the `NeighborConnectivity` declaration the table is bound to. It determines + the table's `domain` -- the declaration's domain extended by its local dimension -- and its + `codomain`. A table alone cannot name its declaration: one sharing another connectivity's + local dimension has a table over the same domain as the owner's, with another codomain. So the + record is built where a table is bound, from its offset-provider key (`offset_provider_to_type`, + `check_neighbor_table`), or given directly for ahead-of-time compilation. + + A table bound under a name that no declaration answers to, as hand-written IR binds them, has + no declaration: `connectivity` is then the table's own structural `ConnectivityType`, which is + also what `NeighborTable.__gt_type__()` returns. + """ + + connectivity: type[NeighborConnectivity] | ConnectivityType + dtype: core_defs.DType + skip_value: Optional[core_defs.IntegralScalar] + #: The table's number of entries per element. A declaration may leave it to the table; where + #: it states one, `check_neighbor_table` checks the table against it. max_neighbors: int @property - def source_dim(self) -> Dimension: - return self.domain[0] + def domain(self) -> tuple[Dimension, Dimension]: + if isinstance(self.connectivity, ConnectivityType): + first, second = self.connectivity.domain + return (first, second) + return (self.connectivity.domain, local_dimension_of(self.connectivity)) @property - def neighbor_dim(self) -> Dimension: - return self.domain[1] + def codomain(self) -> Dimension: + return self.connectivity.codomain + + @property + def has_skip_values(self) -> bool: + return self.skip_value is not None @runtime_checkable @@ -1327,21 +1381,14 @@ def codomain(self) -> DimT_co: """ def __gt_type__(self) -> ConnectivityType: - if is_neighbor_table(self): - return NeighborConnectivityType( - domain=self.domain.dims, - codomain=self.codomain, - dtype=self.dtype, - skip_value=self.skip_value, - max_neighbors=self.ndarray.shape[1], - ) - else: - return ConnectivityType( - domain=self.domain.dims, - codomain=self.codomain, - dtype=self.dtype, - skip_value=self.skip_value, - ) + # NOTE: structural, also for a neighbor table: the table cannot tell which declaration it + # is bound to, so its `NeighborTableType` is built from its offset-provider key. + return ConnectivityType( + domain=self.domain.dims, + codomain=self.codomain, + dtype=self.dtype, + skip_value=self.skip_value, + ) @abc.abstractmethod def inverse_image(self, image_range: UnitRange | NamedRange) -> Sequence[NamedRange]: ... @@ -1472,8 +1519,7 @@ def _connectivity( @runtime_checkable class NeighborTable(Connectivity, Protocol): - # TODO(havogt): work towards encoding this properly in the type - def __gt_type__(self) -> NeighborConnectivityType: ... + def __gt_type__(self) -> ConnectivityType: ... @property def ndarray(self) -> core_defs.NDArrayObject: @@ -1490,16 +1536,17 @@ def is_neighbor_table(obj: Any) -> TypeGuard[NeighborTable]: return ( len(domain_dims) == 2 and domain_dims[0].kind is DimensionKind.HORIZONTAL - and domain_dims[1].kind is DimensionKind.LOCAL + and is_local_dimension(domain_dims[1]) ) OffsetProviderElem: TypeAlias = NeighborTable -OffsetProviderTypeElem: TypeAlias = NeighborConnectivityType -# Note: `OffsetProvider` and `OffsetProviderType` should not be accessed directly, +# Note: `OffsetProvider` and `TableTypes` should not be accessed directly, # use the `get_offset` and `get_offset_type` functions instead. OffsetProvider: TypeAlias = Mapping[Tag, OffsetProviderElem] -OffsetProviderType: TypeAlias = Mapping[Tag, OffsetProviderTypeElem] +#: The types of an offset provider's tables, under the same keys: what transformations and code +#: generation see instead of the tables (ADR 0019). +TableTypes: TypeAlias = Mapping[Tag, NeighborTableType] def is_offset_provider(obj: Any) -> TypeGuard[OffsetProvider]: @@ -1508,25 +1555,54 @@ def is_offset_provider(obj: Any) -> TypeGuard[OffsetProvider]: return all(isinstance(el, OffsetProviderElem) for el in obj.values()) -def is_offset_provider_type(obj: Any) -> TypeGuard[OffsetProviderType]: +def is_table_types(obj: Any) -> TypeGuard[TableTypes]: if not isinstance(obj, Mapping): return False - return all(isinstance(el, OffsetProviderTypeElem) for el in obj.values()) + return all(isinstance(el, NeighborTableType) for el in obj.values()) -def offset_provider_to_type( - offset_provider: OffsetProvider | OffsetProviderType, -) -> OffsetProviderType: +def offset_provider_to_type(offset_provider: OffsetProvider | TableTypes) -> TableTypes: + """The types of an offset provider's tables, each typed by the declaration its key names.""" return { - k: v.__gt_type__() if isinstance(v, Connectivity) else v for k, v in offset_provider.items() + key: value if isinstance(value, NeighborTableType) else _neighbor_table_type(key, value) + for key, value in offset_provider.items() } +def _unbound_table_type(table: NeighborTable) -> NeighborTableType: + structure = table.__gt_type__() + return NeighborTableType( + connectivity=structure, + dtype=structure.dtype, + skip_value=structure.skip_value, + max_neighbors=len(table.domain[1].unit_range), + ) + + +def _neighbor_table_type(key: Tag, table: NeighborTable) -> NeighborTableType: + """ + The type of `table` bound under `key`: typed by the declaration `key` is the `offset_tag` of. + + The declaration is found through the table's local dimension, which knows its owner and the + connectivities sharing it; a key none of them answers to leaves the table undeclared. + """ + table_type = _unbound_table_type(table) + local = table_type.domain[1] + if not (isinstance(local, DimensionMeta) and issubclass(local, LocalDimensionIndex)): + return table_type + # NOTE: the most recent sharer first: a redefined declaration (a re-run notebook cell) is + # appended again under the same tag. + for connectivity in (local.owner, *reversed(local.sharers)): + if connectivity is not None and connectivity.offset_tag == key: + return check_neighbor_table(connectivity, table_type) + return table_type + + def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProviderElem: """ - Get the `OffsetProviderElem` or `OffsetProviderTypeElem` for the given `offset` string. + Get the `OffsetProviderElem` or `NeighborTableType` for the given `offset` string. - Note: All accesses of `OffsetProvider` or `OffsetProviderType` should go through this function. + Note: All accesses of `OffsetProvider` or `TableTypes` should go through this function. """ # TODO(havogt): Once we have a custom class for `OffsetProvider`, we can absorb this functionality into it. if offset_tag not in offset_provider: @@ -1534,10 +1610,52 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid return offset_provider[offset_tag] # TODO return a valid dimension -get_offset_type: Callable[[OffsetProviderType, str], OffsetProviderTypeElem] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and OffsetProviderType overlap +get_offset_type: Callable[[TableTypes, str], NeighborTableType] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and TableTypes overlap + + +def connectivity_key_over( + offset_provider: OffsetProvider | TableTypes, local_dim: Dimension | Tag +) -> str: + """ + 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 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 `NeighborConnectivity`), so which one does + not matter. + + Raises: + KeyError: If no bound connectivity has `local_dim` as its local dimension. + """ + 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 has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: str) -> bool: +def _neighbor_dim_of(connectivity: Any) -> Optional[Dimension]: + if isinstance(connectivity, NeighborTableType): + return connectivity.domain[1] + if is_neighbor_table(connectivity): + return connectivity.domain.dims[1] + return None + + +def has_offset(offset_provider: OffsetProvider | TableTypes, offset_tag: str) -> bool: """Determine if offset provider has an element for the given offset tag.""" try: get_offset(offset_provider, offset_tag) # type: ignore[arg-type] # implementation is shared with `get_offset_type`, no need to duplicate the function @@ -1637,8 +1755,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() @@ -1658,8 +1776,8 @@ class GridType(StrEnum): def order_dimensions(dims: Iterable[Dimension]) -> list[Dimension]: """Find the canonical ordering of the dimensions in `dims`.""" - if sum(1 for dim in dims if dim.kind == DimensionKind.LOCAL) > 1: - raise ValueError("There are more than one dimension with DimensionKind 'LOCAL'.") + if sum(1 for dim in dims if is_local_dimension(dim)) > 1: + raise ValueError("There is more than one local dimension.") # NOTE: `__qualname__`, not `tag`. The tag is qualified, so ordering by it would make a # field's canonical dimension order depend on *which module* each dimension is declared in -- # moving a declaration would silently reorder a field's dimensions. The unqualified name keeps @@ -1668,7 +1786,7 @@ def order_dimensions(dims: Iterable[Dimension]) -> list[Dimension]: return sorted( dims, key=lambda dim: ( - _DIM_KIND_ORDER[dim.kind], + _dimension_rank(dim), as_non_staggered(dim).__qualname__, as_non_staggered(dim).tag, ), @@ -1699,7 +1817,7 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: Find an ordering of multiple lists of dimensions. The resulting list contains all unique dimensions from the input lists, - sorted first by dims_kind_order, i.e., `Dimension.kind` (`HORIZONTAL` < `LOCAL` < `VERTICAL`) and then + sorted first by horizontal < local < vertical (see `order_dimensions`) and then lexicographically by `Dimension.tag`. Examples: @@ -1707,8 +1825,8 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: >>> class I(CartesianAxisIndex, kind=DimensionKind.HORIZONTAL): ... >>> class J(CartesianAxisIndex, kind=DimensionKind.HORIZONTAL): ... >>> class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... - >>> class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... - >>> class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... + >>> class E2V(LocalDimensionIndex): ... + >>> class E2C(LocalDimensionIndex): ... >>> promote_dims([J, K], [I, K]) == [I, J, K] True >>> promote_dims([K, J], [I, K]) @@ -1720,7 +1838,7 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: >>> promote_dims([I, E2C], [E2V, K]) Traceback (most recent call last): ... - ValueError: There are more than one dimension with DimensionKind 'LOCAL'. + ValueError: There is more than one local dimension. """ for dims in dims_list: @@ -1878,20 +1996,6 @@ def __init_subclass__(cls, /, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) -class ConstList(DimensionIndex, kind=DimensionKind.LOCAL): - """ - The local dimension of a list of one repeated value (`make_const_list`). - - The value is broadcast against the neighbor lists it is combined with, and a materialized - constant list has extent 1 along it. It indexes no table, so it is never in an offset provider. - - Declared here, once: it used to be built independently in `iterator/embedded.py` and in the - DaCe lowering, which only worked while dimensions compared by `(name, kind)`. - """ - - __slots__ = () - - def _reduce_staggered(cls: StaggeredMeta) -> Any: """ Pickle a staggered dimension through its base, falling back to by-reference. @@ -1961,3 +2065,391 @@ 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): + """ + 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.owner, LsqCoeff.max_neighbors, str(LsqCoeff) + (None, 3, 'LsqCoeff[local]') + + A local dimension has no `kind`: it is `None`, and `is_local_dimension` reads localness + from the class. + + Neighbor counts are optional. A declared count is a constraint the bound table has to + satisfy (see `check_neighbor_table`); an undeclared one is taken from the table. + """ + + __slots__ = () + + kind: ClassVar[Optional[DimensionKind]] = None + + #: The connectivity this dimension is the local axis of, or `None` if it indexes no table. + #: Set by `NeighborConnectivity` when the connectivity is declared. + owner: ClassVar[Optional[type[NeighborConnectivity]]] = None + #: The connectivities sharing this dimension with its owner, in declaration order. Kept so that + #: a table bound under a sharer's `offset_tag` can be typed by its declaration: the table alone + #: looks like the owner's (see `NeighborTableType`). + sharers: ClassVar[tuple[type[NeighborConnectivity], ...]] = () + #: Number of entries per element, i.e. the table's second extent, if declared. + max_neighbors: ClassVar[Optional[int]] = None + #: Least number of *valid* neighbors of any element, if declared. Fewer than + #: `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, + /, + *, + size: Optional[int] = None, + kind: Optional[DimensionKind] = None, + **kwargs: Any, + ) -> None: + if kind is not None: + raise TypeError( + f"'{cls.__qualname__}' is a local dimension and has no kind; got '{kind}'." + ) + super().__init_subclass__(**kwargs) + # NOTE: reset rather than inherited: a subclass of an owned local dimension is a + # different dimension, and does not index its parent's table. + cls.owner = None + cls.sharers = () + cls.declared_size = _check_neighbor_count(cls, "size", size) + cls.max_neighbors = cls.min_neighbors = cls.declared_size + + +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): + 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) + + +class ConstList(LocalDimensionIndex, size=1): + """ + The local dimension of a list of one repeated value (`make_const_list`). + + An owner-less local dimension of size 1: the value is broadcast against the neighbor lists it + is combined with, and a materialized constant list has extent 1 along it. It indexes no table, + so it is never in an offset provider. + + Declared here, once: it used to be built independently in `iterator/embedded.py` and in the + DaCe lowering, which only worked while dimensions compared by `(name, kind)`. + """ + + __slots__ = () + + +class ConnectivityMeta(type): + """ + Metaclass of `NeighborConnectivity` declarations. + + 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. + """ + + # 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`. + domain: Dimension + codomain: Dimension + + @property + 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 := cls.__dict__.get("Local")) is None: + raise TypeError( + f"'{cls.__qualname__}' is not a connectivity declaration; declare one by" + " subclassing 'NeighborConnectivity[Domain, Codomain]'." + ) + return cast(type[LocalDimensionIndex], local) + + 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)] + 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) + + 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, tagged with `offset_tag`. + """ + 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.domain, cls._local()) + ) + type.__setattr__(cls, "_field_offset", field_offset) + return field_offset + + +class NeighborConnectivity[Domain: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + """ + Declare a neighbor connectivity: for each `Domain` element, a list of `Codomain` neighbors. + + The declaration names the connectivity's local dimension -- its nested `Local` class -- + and optionally its neighbor counts. It holds no data: the neighbor table is bound at call + time through the offset provider. `check_neighbor_table` checks a table 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.domain is Vertex, V2E.codomain is Edge + (True, True) + >>> V2E.Local.owner is V2E, V2E.Local.max_neighbors, V2E.Local.min_neighbors + (True, 6, 5) + """ + + # NOTE: `Local` is not annotated (see `ConnectivityMeta`); every subclass declares it, as a + # nested class or as `Local: TypeAlias = `. + domain: 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 = [ + get_args(base) + for base in cls.__dict__.get("__orig_bases__", ()) + if get_origin(base) is NeighborConnectivity + ] + if len(params) != 1 or len(params[0]) != 2: + raise TypeError( + f"'{name}' must derive from 'NeighborConnectivity[Domain, Codomain]' directly," + " with both dimensions given." + ) + domain, codomain = params[0] + for role, dim in (("Domain", domain), ("Codomain", codomain)): + if not isinstance(dim, DimensionMeta) or is_local_dimension(dim): + raise TypeError(f"'{name}': '{role}' must be a non-local dimension, got '{dim}'.") + + local = cls.__dict__.get("Local") + if not (isinstance(local, DimensionMeta) and issubclass(local, LocalDimensionIndex)): + raise TypeError( + f"'{name}' must declare its local dimension, either as a nested class" + " ('class Local(LocalDimensionIndex): ...') or by adopting one" + " ('Local: TypeAlias = SomeLocalDim')." + ) + if local is ConstList: + raise TypeError( + f"'{name}' cannot adopt '{ConstList.__qualname__}': it is the local dimension" + " of 'make_const_list' results and belongs to no connectivity." + ) + max_neighbors = _check_neighbor_count(cls, "max_neighbors", max_neighbors) + min_neighbors = _check_neighbor_count(cls, "min_neighbors", min_neighbors) + # NOTE: a declaration whose tag is the owner's is a *redefinition* of it (a re-run + # 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. + owner_name = local.owner.__qualname__ + if domain is not local.owner.domain: + raise TypeError( + f"'{name}' cannot share the local dimension of '{owner_name}': it has domain" + f" '{domain}', but the neighbors of '{owner_name}' are those of" + f" '{local.owner.domain}'." + ) + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("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.domain, cls.codomain = domain, codomain + local.sharers = (*local.sharers, cls) + return + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("min_neighbors", min_neighbors), + ): + # 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__}' ({local.declared_size})." + ) + 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 + and min_neighbors > max_neighbors + ): + raise TypeError( + f"'{name}': 'min_neighbors' ({min_neighbors}) exceeds 'max_neighbors'" + f" ({max_neighbors})." + ) + + cls.domain, cls.codomain = domain, codomain + local.owner = cls + 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 | NeighborTableType, +) -> NeighborTableType: + """ + 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). + + Returns: + The type of the table bound to `connectivity`. + + Raises: + ValueError: On the first mismatch, naming the connectivity and the mismatch. + """ + name = connectivity.__qualname__ + local = local_dimension_of(connectivity) + + def fail(reason: str) -> NoReturn: + raise ValueError(f"The table bound to '{name}' does not match its declaration: {reason}.") + + if isinstance(table, NeighborTableType): + table_type = table + elif is_neighbor_table(table): + table_type = _unbound_table_type(table) + else: + fail(f"expected a neighbor table, got '{table}'") + if isinstance(table_type.connectivity, ConnectivityMeta) and ( + table_type.connectivity is not connectivity + ): + fail(f"its type is bound to '{table_type.connectivity.__qualname__}'") + expected_domain = (connectivity.domain, local) + if tuple(table_type.domain) != expected_domain: + fail( + f"its domain is '({', '.join(map(str, table_type.domain))})'," + 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 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( + 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: + 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," + " 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}" + ) + return dataclasses.replace(table_type, connectivity=connectivity) diff --git a/src/gt4py/next/constructors.py b/src/gt4py/next/constructors.py index 2bef28dce0..201b686764 100644 --- a/src/gt4py/next/constructors.py +++ b/src/gt4py/next/constructors.py @@ -653,7 +653,7 @@ def as_connectivity( >>> from gt4py import next as gtx >>> class Vertex(gtx.DimensionIndex): ... >>> class Edge(gtx.DimensionIndex): ... - >>> class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + >>> class V2EDim(gtx.LocalDimensionIndex): ... >>> data = np.array([[0, 1], [1, 2], [2, 0]]) >>> conn = gtx.as_connectivity([Vertex, V2EDim], Edge, data) >>> conn.ndarray diff --git a/src/gt4py/next/custom_layout_allocators.py b/src/gt4py/next/custom_layout_allocators.py index 5eee3833ac..453aee0bad 100644 --- a/src/gt4py/next/custom_layout_allocators.py +++ b/src/gt4py/next/custom_layout_allocators.py @@ -160,7 +160,7 @@ def pos_of_kind(kind: common.DimensionKind) -> list[int]: horizontals = pos_of_kind(common.DimensionKind.HORIZONTAL) verticals = pos_of_kind(common.DimensionKind.VERTICAL) - locals_ = pos_of_kind(common.DimensionKind.LOCAL) + locals_ = [i for i, dim in enumerate(dims) if common.is_local_dimension(dim)] layout_map = [0] * len(dims) for i, pos in enumerate(horizontals + verticals + locals_): diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index b8133cbafa..26504c332d 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -230,7 +230,9 @@ def as_scalar(self) -> core_defs.ScalarT: def premap( self: NdArrayField, - *connectivities: common.Connectivity | fbuiltins.FieldOffset, + *connectivities: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], ) -> NdArrayField: """ Rearrange the field content using the provided connectivities (index mappings). @@ -305,7 +307,10 @@ def premap( codomains_counter: collections.Counter[common.Dimension] = collections.Counter() for connectivity in connectivities: - # For neighbor reductions, a FieldOffset is passed instead of an actual Connectivity + # For neighbor reductions, a FieldOffset or a connectivity declaration is passed + # instead of an actual Connectivity + if isinstance(connectivity, common.ConnectivityMeta): + connectivity = connectivity.__gt_field_offset__() if not isinstance(connectivity, common.Connectivity): assert isinstance(connectivity, fbuiltins.FieldOffset) connectivity = connectivity.as_connectivity_field() @@ -357,8 +362,10 @@ def premap( def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: return functools.reduce( lambda field, current_index_field: field.premap(current_index_field), @@ -960,11 +967,11 @@ def _builtin_op( ) -> NdArrayField[common.DimsT, core_defs.ScalarT]: xp = field.array_ns - if not axis.kind == common.DimensionKind.LOCAL: + if not common.is_local_dimension(axis): raise ValueError("Can only reduce local dimensions.") if axis not in field.domain.dims: raise ValueError(f"Field can not be reduced as it doesn't have dimension '{axis}'.") - if len([d for d in field.domain.dims if d.kind is common.DimensionKind.LOCAL]) > 1: + if len([d for d in field.domain.dims if common.is_local_dimension(d)]) > 1: raise NotImplementedError( "Reducing a field with more than one local dimension is not supported." ) @@ -972,8 +979,8 @@ def _builtin_op( current_offset_provider = embedded_context.get_offset_provider(None) assert current_offset_provider is not None offset_definition = common.get_offset( - current_offset_provider, axis.tag - ) # assumes offset and local dimension have same name + current_offset_provider, common.connectivity_key_over(current_offset_provider, axis) + ) assert common.is_neighbor_table(offset_definition) new_domain = common.Domain(*[nr for nr in field.domain if nr.dim != axis]) diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index 1428e664d1..a464de9763 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -159,9 +159,9 @@ def _make_compiled_programs_pool( def compile( self, - offset_provider: common.OffsetProviderType + offset_provider: common.TableTypes | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + | list[common.TableTypes | common.OffsetProvider] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -185,9 +185,7 @@ def compile( ) if self.compilation_options.connectivities is None and offset_provider is None: - raise ValueError( - "Cannot compile a program without connectivities / OffsetProviderType." - ) + raise ValueError("Cannot compile a program without connectivities / TableTypes.") if not all(isinstance(v, list) for v in static_args.values()): raise TypeError( "Please provide the static arguments as lists." @@ -200,8 +198,7 @@ def compile( offset_provider = [offset_provider] # type: ignore[list-item] # cleanup offset_provider vs offset_provider_type assert all( - common.is_offset_provider(op) or common.is_offset_provider_type(op) - for op in offset_provider + common.is_offset_provider(op) or common.is_table_types(op) for op in offset_provider ) self._compiled_programs.compile(offset_providers=offset_provider, **static_args) @@ -487,9 +484,9 @@ def __call__( @override def compile( self, - offset_provider: common.OffsetProviderType + offset_provider: common.TableTypes | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + | list[common.TableTypes | common.OffsetProvider] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 7474cd1406..dc8552425c 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -130,9 +130,9 @@ def _type_conversion_helper(t: type) -> type[ts.TypeSpec] | tuple[type[ts.TypeSp elif t is common.Dimension: return ts.DimensionType elif t is FieldOffset: - return ts.OffsetType + return ts.ShiftType elif t is common.Connectivity: - return ts.OffsetType + return ts.ShiftType elif t is core_defs.ScalarT: return ts.ScalarType elif t is common.Domain: @@ -490,11 +490,20 @@ def _cache(self) -> dict: return {} def __post_init__(self) -> None: - if len(self.target) == 2 and self.target[1].kind != common.DimensionKind.LOCAL: + if len(self.target) == 2 and not common.is_local_dimension(self.target[1]): raise ValueError("Second dimension in offset must be a local dimension.") - def __gt_type__(self) -> ts.OffsetType: - return ts.OffsetType(source=self.source, target=self.target, tag=self.value) + def __gt_type__(self) -> ts.ShiftType: + return ts.ShiftType(codomain=self.source, domain=self.target, tag=self.value) + + @property + def Local(self) -> common.Dimension: + """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.""" @@ -532,10 +541,11 @@ def as_connectivity_field(self) -> common.Connectivity: return connectivity -def is_cartesian_offset(offset: FieldOffset | ts.OffsetType) -> bool: +def is_cartesian_offset(offset: FieldOffset | ts.ShiftType) -> bool: + shift_type = offset.__gt_type__() if isinstance(offset, FieldOffset) else offset return ( - len(offset.target) == 1 - and offset.source == offset.target[0] - and offset.source.kind == offset.target[0].kind - and offset.target[0].kind != common.DimensionKind.LOCAL + len(shift_type.domain) == 1 + and shift_type.codomain == shift_type.domain[0] + and shift_type.codomain.kind == shift_type.domain[0].kind + and not common.is_local_dimension(shift_type.domain[0]) ) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 9e8f1fa0f1..1ef443c17c 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -434,11 +434,23 @@ 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 + # dimension of the domain of the shift it is typed as. + case ts.ShiftType(domain=(_, local)) if node.attr == "Local": + attr_type: ts.TypeSpec = ts.DimensionType(dim=local) + case _: + # NOTE: only attributes that are types themselves: a type's other fields (the + # dimensions of a `ShiftType`, say) are not values in DSL code. + if not isinstance( + type_attr := getattr(new_value.type, node.attr, None), ts.TypeSpec + ): + raise errors.DSLError( + node.location, f"'{new_value.type}' has no attribute '{node.attr}'." + ) + attr_type = type_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: @@ -456,20 +468,21 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri f"Tuples need to be indexed with literal integers, got '{node.index}'.", ) from ex new_type = types[index] - case ts.OffsetType(source=source, target=(target1, target2), tag=tag): - if not target2.kind == DimensionKind.LOCAL: + case ts.ShiftType(codomain=codomain, domain=(domain, local), tag=tag): + if not common.is_local_dimension(local): raise errors.DSLError( new_value.location, "Second dimension in offset must be a local dimension." ) - new_type = ts.OffsetType(source=source, target=(target1,), tag=tag) - case ts.OffsetType(source=source, target=(target,), tag=tag): + new_type = ts.ShiftType(codomain=codomain, domain=(domain,), tag=tag) + case ts.ShiftType(codomain=codomain, domain=(domain,), tag=tag): # for cartesian axes (e.g. I, J) the index of the subscript only # signifies the displacement in the respective dimension, - # but does not change the target type. - if source != target: + # but does not change the domain. + if codomain != domain: raise errors.DSLError( new_value.location, - "Source and target must be equal for offsets with a single target.", + "Codomain and domain must be equal for a shift with a single domain" + " dimension.", ) if tag is None: raise errors.DSLError( @@ -483,7 +496,7 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri ) ], hints=[ - f"Write the displacement directly, e.g. '{source.__qualname__} + 1'." + f"Write the displacement directly, e.g. '{codomain.__qualname__} + 1'." ], ) new_type = new_value.type @@ -625,7 +638,7 @@ def _deduce_compare_type( def _deduce_binop_type( self, node: foast.BinOp, *, left: foast.Expr, right: foast.Expr, **kwargs: Any ) -> Optional[ts.TypeSpec]: - if isinstance(left.type, ts.OffsetType): + if isinstance(left.type, ts.ShiftType): raise errors.DSLError( node.location, f"Type '{left.type}' can not be used in operator '{node.op}'." ) @@ -733,7 +746,7 @@ def _deduce_binop_type( ], ) conn = common.connectivity_for_cartesian_shift(left.type.dim, offset_index) - return ts.OffsetType(source=conn.codomain, target=(conn.domain_dim,)) + return ts.ShiftType(codomain=conn.codomain, domain=(conn.domain_dim,)) else: raise errors.DSLError(node.location, err_msg) @@ -802,8 +815,8 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call: # meaningful unsubscripted, as the neighbor access `field(Off)`. if ( isinstance(arg, (foast.Name, foast.Attribute)) - and isinstance(arg.type, ts.OffsetType) - and len(arg.type.target) == 1 + and isinstance(arg.type, ts.ShiftType) + and len(arg.type.domain) == 1 ): raise errors.DSLError( arg.location, @@ -811,7 +824,7 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call: hints=[f"Give the displacement, e.g. '{arg!s}[1]'."], ) elif isinstance(new_func.type, ts.DimensionType): - assert new_func.type.dim.kind == DimensionKind.LOCAL + assert common.is_local_dimension(new_func.type.dim) return foast.Call( func=new_func, args=new_args, @@ -1000,15 +1013,15 @@ def _visit_astype(self, node: foast.Call, **kwargs: Any) -> foast.Call: def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: arg_0 = node.args[0].type arg_1 = node.args[1].type - assert isinstance(arg_0, ts.OffsetType) + assert isinstance(arg_0, ts.ShiftType) assert isinstance(arg_1, ts.FieldType) if not fbuiltins.is_cartesian_offset(arg_0): - target_dims = ", ".join(d.__qualname__ for d in arg_0.target) # for the diagnostic + domain_dims = ", ".join(d.__qualname__ for d in arg_0.domain) # for the diagnostic raise errors.DSLError( node.location, f"'as_offset' is only supported for Cartesian offsets " - f"(single target dimension equal to source dimension); " - f"got source '{arg_0.source.__qualname__}' and target ({target_dims}).", + f"(a single domain dimension equal to the codomain); " + f"got codomain '{arg_0.codomain.__qualname__}' and domain ({domain_dims}).", ) if not type_info.is_integral(arg_1): raise errors.DSLError( @@ -1018,11 +1031,11 @@ def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: f"{node.location}", ) - if arg_0.source not in arg_1.dims: + if arg_0.codomain not in arg_1.dims: raise errors.DSLError( node.location, f"Incompatible argument in call to '{node.func!s}': " - f"'{arg_0.source}' not in list of offset field dimensions '{arg_1.dims}'. " + f"'{arg_0.codomain}' not in list of offset field dimensions '{arg_1.dims}'. " f"{node.location}", ) diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index f4c8a4fb10..aa6731e11c 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -297,7 +297,7 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: # `field(Off[idx])` # (matched on the type, not the node, to also accept `mod.Off[idx]`) case foast.Subscript( - value=foast.LocatedNode(type=ts.OffsetType(tag=str() as offset_tag)), + value=foast.LocatedNode(type=ts.ShiftType(tag=str() as offset_tag)), index=index, ): # Constant folding to a `Literal` ensures that `index` becomes an `OffsetLiteral`, @@ -331,8 +331,8 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: case foast.Call(func=foast.Name(id="as_offset")): func_args = arg offset_type = func_args.args[0].type - assert isinstance(offset_type, ts.OffsetType) - dim = offset_type.source + assert isinstance(offset_type, ts.ShiftType) + dim = offset_type.codomain offset_field = self.visit(func_args.args[1], **kwargs) current_expr = im.as_fieldop( im.lambda_("__it", "__offset")( @@ -342,7 +342,7 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: ) )(current_expr, offset_field) # `field(Off)` - case foast.LocatedNode(type=ts.OffsetType(tag=str() as offset_tag, target=(_, _))): + case foast.LocatedNode(type=ts.ShiftType(tag=str() as offset_tag, domain=(_, _))): # only a single unstructured shift is supported so returning here is fine even though we # are in a loop. assert len(node.args) == 1 diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 32b9f9dfac..595ac45afd 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() @@ -406,7 +406,7 @@ def _construct_itir_domain_arg( dim_stop, ) - if dim.kind == common.DimensionKind.LOCAL: + if common.is_local_dimension(dim): raise ValueError(f"common.Dimension '{dim.__qualname__}' must not be local.") domain_args.append( itir.FunCall( diff --git a/src/gt4py/next/ffront/transform_utils.py b/src/gt4py/next/ffront/transform_utils.py index 09c9d4b9ee..2b441796a6 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,10 +61,12 @@ 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: + if isinstance(o, common.DimensionMeta) and common.is_local_dimension(o): deduced_grid_type = common.GridType.UNSTRUCTURED break diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index 52a7081b1d..b7cac93ade 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -231,6 +231,24 @@ 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; 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.domain, + obj.codomain, + 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__ + 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/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 7b37e35684..73c5e10102 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -120,15 +120,6 @@ def __init__( def __gt_origin__(self) -> typing.Never: raise NotImplementedError - def __gt_type__(self) -> common.NeighborConnectivityType: - return common.NeighborConnectivityType( - domain=self.domain_dims, - codomain=self.codomain_dim, - max_neighbors=self._max_neighbors, - skip_value=self.skip_value, - dtype=self.dtype, - ) - @property def domain(self) -> common.Domain: return common.Domain( @@ -151,7 +142,12 @@ def ndarray(self) -> core_defs.NDArrayObject: def asnumpy(self) -> np.ndarray: raise NotImplementedError - def premap(self, index_field: common.Connectivity | fbuiltins.FieldOffset) -> common.Field: + def premap( + self, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + ) -> common.Field: raise NotImplementedError def restrict( # type: ignore[override] @@ -168,8 +164,10 @@ def as_scalar(self) -> typing.Never: def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError() @@ -562,9 +560,14 @@ def execute_shift( if tag == common.ConstList.tag: new_entry[i] = 0 else: - offset_implementation = common.get_offset(offset_provider, tag) + # NOTE: the sparse tag is the local dimension's; the table over it may be + # keyed by a connectivity sharing it (see `common.connectivity_key_over`). + offset_implementation = common.get_offset( + offset_provider, + common.connectivity_key_over(offset_provider, tag), + ) assert common.is_neighbor_table(offset_implementation) - source_dim = offset_implementation.__gt_type__().source_dim + source_dim = offset_implementation.__gt_type__().domain[0] cur_index = pos[source_dim.tag] assert common.is_int_index(cur_index) if offset_implementation[cur_index, index].as_scalar() in [ @@ -586,7 +589,7 @@ def execute_shift( return new_pos offset_implementation = common.get_offset(offset_provider, tag) if common.is_neighbor_table(offset_implementation): - source_dim = offset_implementation.__gt_type__().source_dim + source_dim = offset_implementation.__gt_type__().domain[0] assert source_dim.tag in pos new_pos = pos.copy() new_pos.pop(source_dim.tag) @@ -910,7 +913,7 @@ def _get_sparse_dimensions(axes: Sequence[common.Dimension]) -> list[common.Dime return [ axis for axis in axes - if isinstance(axis, common.DimensionMeta) and axis.kind == common.DimensionKind.LOCAL + if isinstance(axis, common.DimensionMeta) and common.is_local_dimension(axis) ] @@ -991,9 +994,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.offset.value: i}) # type: ignore[dict-item] + self._translate_named_indices({**named_indices, local_tag: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ @@ -1143,8 +1147,10 @@ def as_scalar(self) -> core_defs.IntegralScalar: def premap( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1284,8 +1290,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() @@ -1400,16 +1408,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__().domain[1] @dataclasses.dataclass(frozen=True) @@ -1439,7 +1457,7 @@ def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: return _List( values=tuple( shifted.deref() - for i in range(connectivity.__gt_type__().max_neighbors) + for i in range(len(connectivity.domain[1].unit_range)) if (shifted := it.shift(offset_str, i)).can_deref() ), offset=offset, @@ -1454,12 +1472,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) @@ -1512,17 +1532,20 @@ 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, self.list_offset) + connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( values=tuple( shifted.deref() - for i in range(connectivity.__gt_type__().max_neighbors) + for i in range(len(connectivity.domain[1].unit_range)) if ( shifted := self.it.shift(*self.offsets, SparseTag(self.list_offset), i) ).can_deref() ), - offset=runtime.Offset(value=self.list_offset), + offset=runtime.Offset(value=connectivity_key), ) def can_deref(self) -> bool: @@ -1653,9 +1676,9 @@ def _dimension_to_tag( return {k.tag: v for k, v in domain.items()} -def _validate_domain(domain: Domain, offset_provider_type: common.OffsetProviderType) -> None: +def _validate_domain(domain: Domain, offset_provider_type: common.TableTypes) -> None: if isinstance(domain, runtime.CartesianDomain): - if any(isinstance(o, common.ConnectivityType) for o in offset_provider_type.values()): + if any(isinstance(o, common.NeighborTableType) for o in offset_provider_type.values()): raise RuntimeError( "Got a 'CartesianDomain', but found a 'Connectivity' in 'offset_provider', expected 'UnstructuredDomain'." ) @@ -1769,11 +1792,13 @@ 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), - common.named_range((offset_type, connectivity.__gt_type__().max_neighbors)), + common.named_range((offset_type, len(connectivity.domain[1].unit_range))), ), type_.element_type return domain, type_ diff --git a/src/gt4py/next/iterator/ir.py b/src/gt4py/next/iterator/ir.py index 717188c8a9..bfb0c666ed 100644 --- a/src/gt4py/next/iterator/ir.py +++ b/src/gt4py/next/iterator/ir.py @@ -99,9 +99,10 @@ def dim(self) -> common.Dimension: return common.resolve(self.value) @property - def kind(self) -> common.DimensionKind: + def kind(self) -> Optional[common.DimensionKind]: # NOTE: derived, not stored: the dimension class carries its kind, so a stored copy could - # only disagree with it (it used to, for local dimensions printed as vertical). + # only disagree with it (it used to, for local dimensions printed as vertical). `None` for a + # local dimension, see `common.is_local_dimension`. return self.dim.kind diff --git a/src/gt4py/next/iterator/ir_utils/domain_utils.py b/src/gt4py/next/iterator/ir_utils/domain_utils.py index b23ef3a934..b87f340096 100644 --- a/src/gt4py/next/iterator/ir_utils/domain_utils.py +++ b/src/gt4py/next/iterator/ir_utils/domain_utils.py @@ -172,7 +172,7 @@ def translate( | Literal[trace_shifts.Sentinel.VALUE, trace_shifts.Sentinel.ALL_NEIGHBORS], ..., ], - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, #: A dictionary mapping axes names to their length. See #: func:`gt4py.next.iterator.transforms.infer_domain.infer_expr` for more details. symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, @@ -208,13 +208,13 @@ def translate( trace_shifts.Sentinel.VALUE, ] - connectivity: common.NeighborTable | common.NeighborConnectivityType + connectivity: common.NeighborTable | common.NeighborTableType if common.is_offset_provider(offset_provider): connectivity = common.get_offset(offset_provider, off.value) old_dim = connectivity.domain.dims[0] new_dim = connectivity.codomain else: - assert common.is_offset_provider_type(offset_provider) + assert common.is_table_types(offset_provider) connectivity = common.get_offset_type(offset_provider, off.value) old_dim = connectivity.domain[0] new_dim = connectivity.codomain diff --git a/src/gt4py/next/iterator/pretty_printer.py b/src/gt4py/next/iterator/pretty_printer.py index a7376e0d22..201044177c 100644 --- a/src/gt4py/next/iterator/pretty_printer.py +++ b/src/gt4py/next/iterator/pretty_printer.py @@ -138,8 +138,8 @@ def implied_literal_type(value: str) -> ts.ScalarType: _AXIS_KIND_SUFFIX: Final = { common.DimensionKind.HORIZONTAL: "ₕ", common.DimensionKind.VERTICAL: "ᵥ", - common.DimensionKind.LOCAL: "ₗ", } +_LOCAL_AXIS_SUFFIX: Final = "ₗ" class PrettyPrinter(NodeTranslator): @@ -240,7 +240,13 @@ def visit_AxisLiteral(self, node: ir.AxisLiteral, *, prec: int) -> list[str]: dim: Optional[common.Dimension] = node.type.dim else: dim = common.resolve_loaded(node.value) - kind = _AXIS_KIND_SUFFIX[dim.kind] if dim is not None else "ₕ" + if dim is None: + kind = "ₕ" + elif common.is_local_dimension(dim): + kind = _LOCAL_AXIS_SUFFIX + else: + assert dim.kind is not None + kind = _AXIS_KIND_SUFFIX[dim.kind] return [str(node.value) + kind] def visit_SymRef(self, node: ir.SymRef, *, prec: int) -> list[str]: diff --git a/src/gt4py/next/iterator/runtime.py b/src/gt4py/next/iterator/runtime.py index 88a466229f..3b604fcb31 100644 --- a/src/gt4py/next/iterator/runtime.py +++ b/src/gt4py/next/iterator/runtime.py @@ -133,9 +133,7 @@ def fendef( ) -def _deduce_domain( - domain: dict[common.Dimension, range], offset_provider_type: common.OffsetProviderType -): +def _deduce_domain(domain: dict[common.Dimension, range], offset_provider_type: common.TableTypes): if isinstance(domain, UnstructuredDomain): domain_builtin = builtins.unstructured_domain elif isinstance(domain, CartesianDomain): @@ -143,7 +141,7 @@ def _deduce_domain( else: domain_builtin = ( builtins.unstructured_domain - if any(isinstance(o, common.ConnectivityType) for o in offset_provider_type.values()) + if any(isinstance(o, common.NeighborTableType) for o in offset_provider_type.values()) else builtins.cartesian_domain ) diff --git a/src/gt4py/next/iterator/transforms/collapse_tuple.py b/src/gt4py/next/iterator/transforms/collapse_tuple.py index 08ad5c9ede..1d049897d9 100644 --- a/src/gt4py/next/iterator/transforms/collapse_tuple.py +++ b/src/gt4py/next/iterator/transforms/collapse_tuple.py @@ -187,7 +187,7 @@ def apply( node: itir.Node, *, remove_letified_make_tuple_elements: bool = True, - offset_provider_type: Optional[common.OffsetProviderType] = None, + offset_provider_type: Optional[common.TableTypes] = None, within_stencil: Optional[bool] = None, # manually passing enabled transformations is mostly for allowing separate testing of the modes enabled_transformations: Optional[Transformation] = None, diff --git a/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py b/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py index 40d956fca0..aaf34d5759 100644 --- a/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py +++ b/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py @@ -27,7 +27,7 @@ def apply( cls, node: itir.Node, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, allow_undeclared_symbols: bool = False, ) -> itir.Node: node = type_inference.infer( diff --git a/src/gt4py/next/iterator/transforms/cse.py b/src/gt4py/next/iterator/transforms/cse.py index 983fcafbd1..78d4612a16 100644 --- a/src/gt4py/next/iterator/transforms/cse.py +++ b/src/gt4py/next/iterator/transforms/cse.py @@ -462,7 +462,7 @@ def apply( cls, node: ProgramOrExpr, within_stencil: bool | None = None, - offset_provider_type: common.OffsetProviderType | None = None, + offset_provider_type: common.TableTypes | None = None, *, uids: utils.IDGeneratorPool, ) -> ProgramOrExpr: diff --git a/src/gt4py/next/iterator/transforms/dead_code_elimination.py b/src/gt4py/next/iterator/transforms/dead_code_elimination.py index 1ea906ae98..8a7d84b0f2 100644 --- a/src/gt4py/next/iterator/transforms/dead_code_elimination.py +++ b/src/gt4py/next/iterator/transforms/dead_code_elimination.py @@ -17,7 +17,7 @@ def dead_code_elimination( program: itir.Program, *, uids: utils.IDGeneratorPool, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> itir.Program: """ Perform dead code elimination on a program by simplifying or removing diff --git a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py index 6b705a29dc..7028e0837e 100644 --- a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py +++ b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py @@ -61,7 +61,7 @@ def apply( node: ProgramOrExpr, *, uids: utils.IDGeneratorPool | None, - offset_provider_type: common.OffsetProviderType | None = None, + offset_provider_type: common.TableTypes | None = None, ) -> ProgramOrExpr: if node.type is None: node = itir_inference.infer( diff --git a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py index 3e71508e25..9babe30fe8 100644 --- a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py +++ b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py @@ -123,7 +123,7 @@ def fuse_as_fieldop( expr: itir.Expr, eligible_args: list[bool], *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, enable_cse: bool, uids: utils.IDGeneratorPool, ) -> itir.Expr: @@ -301,7 +301,7 @@ def all(self) -> FuseAsFieldOp.Transformation: enabled_transformations = Transformation.all() uids: utils.IDGeneratorPool - offset_provider_type: common.OffsetProviderType + offset_provider_type: common.TableTypes enable_cse: bool # option to disable is mainly for testing purposes @classmethod @@ -309,7 +309,7 @@ def apply( cls, node: itir.Program, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, uids: utils.IDGeneratorPool, allow_undeclared_symbols=False, within_set_at_expr: Optional[bool] = None, diff --git a/src/gt4py/next/iterator/transforms/global_tmps.py b/src/gt4py/next/iterator/transforms/global_tmps.py index b2eafc1090..8952554479 100644 --- a/src/gt4py/next/iterator/transforms/global_tmps.py +++ b/src/gt4py/next/iterator/transforms/global_tmps.py @@ -311,7 +311,7 @@ def _transform_stmt( def create_global_tmps( program: itir.Program, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, #: A dictionary mapping axes names to their length. See :func:`infer_domain.infer_expr` for #: more details. symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, diff --git a/src/gt4py/next/iterator/transforms/infer_domain.py b/src/gt4py/next/iterator/transforms/infer_domain.py index 5466843476..8df09b0860 100644 --- a/src/gt4py/next/iterator/transforms/infer_domain.py +++ b/src/gt4py/next/iterator/transforms/infer_domain.py @@ -58,7 +58,7 @@ class DomainAccessDescriptor(eve.StrEnum): class InferenceOptions(typing.TypedDict): - offset_provider: common.OffsetProvider | common.OffsetProviderType + offset_provider: common.OffsetProvider | common.TableTypes symbolic_domain_sizes: dict[str, itir.Expr] | None allow_uninferred: bool keep_existing_domains: bool @@ -130,7 +130,7 @@ def _extract_accessed_domains( stencil: itir.Expr, input_ids: list[str], target_domain: NonTupleDomainAccess, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]], ) -> dict[str, NonTupleDomainAccess]: accessed_domains: dict[str, NonTupleDomainAccess] = {} @@ -186,7 +186,7 @@ def _infer_as_fieldop( applied_fieldop: itir.FunCall, target_domain: DomainAccess, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]], allow_uninferred: bool, keep_existing_domains: bool, @@ -445,7 +445,7 @@ def infer_expr( expr: _Expr_T, domain: DomainAccess, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, allow_uninferred: bool = False, keep_existing_domains: bool = False, @@ -573,7 +573,7 @@ def _infer_stmt( def infer_program( program: itir.Program, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, allow_uninferred: bool = False, keep_existing_domains: bool = False, diff --git a/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py b/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py index 66cc3af85a..489d140027 100644 --- a/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py +++ b/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py @@ -33,14 +33,14 @@ def _dynamic_shift_args(node: itir.Expr) -> list[bool] | None: @dataclasses.dataclass class InlineDynamicShifts(eve.NodeTranslator, eve.VisitorWithSymbolTableTrait): - offset_provider_type: common.OffsetProviderType + offset_provider_type: common.TableTypes uids: utils.IDGeneratorPool @classmethod def apply( cls, node: itir.Program, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, uids: utils.IDGeneratorPool, ): return cls(offset_provider_type=offset_provider_type, uids=uids).visit(node) diff --git a/src/gt4py/next/iterator/transforms/inline_scalar.py b/src/gt4py/next/iterator/transforms/inline_scalar.py index b424074b5c..223a484702 100644 --- a/src/gt4py/next/iterator/transforms/inline_scalar.py +++ b/src/gt4py/next/iterator/transforms/inline_scalar.py @@ -19,7 +19,7 @@ class InlineScalar(eve.NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) @classmethod - def apply(cls, program: itir.Program, offset_provider_type: common.OffsetProviderType): + def apply(cls, program: itir.Program, offset_provider_type: common.TableTypes): program = itir_inference.infer(program, offset_provider_type=offset_provider_type) return cls().visit(program) diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index ec4363aba3..0865ffda00 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -55,7 +55,7 @@ def _max_domain_range_sizes(offset_provider: common.OffsetProvider) -> dict[str, sizes: dict[str, int] = {} for provider in offset_provider.values(): if common.is_neighbor_table(provider): - src_dim = provider.__gt_type__().source_dim.tag + src_dim = provider.__gt_type__().domain[0].tag codomain_dim = provider.__gt_type__().codomain.tag sizes[src_dim] = max(sizes.get(src_dim, 0), provider.ndarray.shape[0]) sizes[codomain_dim] = max( @@ -134,7 +134,7 @@ def _process_symbolic_domains_option( def apply_common_transforms( ir: itir.Program, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, extract_temporaries=False, unroll_reduce=False, common_subexpression_elimination=True, @@ -147,7 +147,7 @@ def apply_common_transforms( use_max_domain_range_on_unstructured_shift: Optional[bool] = None, ) -> itir.Program: assert isinstance(ir, itir.Program) - # TODO(tehrengruber): Allow `common.OffsetProviderType`, but domain inference currently + # TODO(tehrengruber): Allow `common.TableTypes`, but domain inference currently # relies on static information or `symbolic_domain_sizes`. assert common.is_offset_provider(offset_provider) diff --git a/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py index b9ade1d636..4dd3f598f8 100644 --- a/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py +++ b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py @@ -35,7 +35,7 @@ def _broadcast_to(expr: itir.Expr, target_dims: list[common.Dimension]) -> itir. def _concat_where_with_explicit_broadcast( node: itir.FunCall, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, ) -> itir.FunCall: """ @@ -107,7 +107,7 @@ class _PruneEmptyConcatWhere(PreserveLocationVisitor, NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) - offset_provider: common.OffsetProvider | common.OffsetProviderType + offset_provider: common.OffsetProvider | common.TableTypes symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None @classmethod @@ -115,7 +115,7 @@ def apply( cls: type[Self], node: PRG, *, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, symbolic_domain_sizes: Optional[dict[str, itir.Expr]] = None, ) -> PRG: return cls( diff --git a/src/gt4py/next/iterator/transforms/unroll_reduce.py b/src/gt4py/next/iterator/transforms/unroll_reduce.py index cfb7bb3226..79067bf663 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 ] @@ -52,16 +52,18 @@ def _get_partial_offset_tags(reduce_args: Iterable[itir.Expr]) -> Iterable[str]: def _get_connectivity( applied_reduce_node: itir.FunCall, - offset_provider_type: common.OffsetProviderType, -) -> common.NeighborConnectivityType: + offset_provider_type: common.TableTypes, +) -> common.NeighborTableType: """Return single connectivity that is compatible with the arguments of the reduce.""" if not cpm.is_applied_reduce(applied_reduce_node): raise ValueError("Expected a call to a 'reduce' object, i.e. 'reduce(...)(...)'.") - connectivities: list[common.NeighborConnectivityType] = [] - for o in _get_partial_offset_tags(applied_reduce_node.args): - conn = common.get_offset_type(offset_provider_type, o) - assert isinstance(conn, common.NeighborConnectivityType) + connectivities: list[common.NeighborTableType] = [] + for local_dim in _get_partial_local_dims(applied_reduce_node.args): + conn = common.get_offset_type( + offset_provider_type, common.connectivity_key_over(offset_provider_type, local_dim) + ) + assert isinstance(conn, common.NeighborTableType) connectivities.append(conn) if not connectivities: @@ -85,13 +87,13 @@ class UnrollReduce(PreserveLocationVisitor, NodeTranslator): def apply( cls, node: itir.Node, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, uids: utils.IDGeneratorPool, ) -> itir.Node: return cls(uids=uids).visit(node, offset_provider_type=offset_provider_type) def _visit_reduce( - self, node: itir.FunCall, offset_provider_type: common.OffsetProviderType + self, node: itir.FunCall, offset_provider_type: common.TableTypes ) -> itir.Expr: connectivity_type = _get_connectivity(node, offset_provider_type) max_neighbors = connectivity_type.max_neighbors diff --git a/src/gt4py/next/iterator/type_system/inference.py b/src/gt4py/next/iterator/type_system/inference.py index b878640f12..77d0b033a8 100644 --- a/src/gt4py/next/iterator/type_system/inference.py +++ b/src/gt4py/next/iterator/type_system/inference.py @@ -182,7 +182,7 @@ def on_type_ready(self, cb: Callable[[ts.TypeSpec], None]) -> None: def __call__( self, *args: type_synthesizer.TypeOrTypeSynthesizer, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, **kwargs, ) -> Union[ts.TypeSpec, ObservableTypeSynthesizer]: assert all(isinstance(arg, (ts.TypeSpec, ObservableTypeSynthesizer)) for arg in args), ( @@ -256,7 +256,7 @@ class ITIRTypeInference(eve.NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) - offset_provider_type: Optional[common.OffsetProviderType] + offset_provider_type: Optional[common.TableTypes] #: Allow sym refs to symbols that have not been declared. Mostly used in testing. allow_undeclared_symbols: bool #: Reinference-mode skipping already typed nodes. @@ -267,7 +267,7 @@ def apply( cls, node: T, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, inplace: bool = False, allow_undeclared_symbols: bool = False, ) -> T: @@ -352,7 +352,7 @@ def apply( @classmethod def apply_reinfer( - cls, node: T, *, offset_provider_type: Optional[common.OffsetProviderType] = None + cls, node: T, *, offset_provider_type: Optional[common.TableTypes] = None ) -> T: """ Given a partially typed node infer the type of ``node`` and its sub-nodes. diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 278683429e..910dd68dcb 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -68,7 +68,7 @@ def __post_init__(self): def __call__( self, *args: TypeOrTypeSynthesizer, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, **kwargs, ) -> TypeOrTypeSynthesizer: return self.type_synthesizer(*args, offset_provider_type=offset_provider_type, **kwargs) @@ -313,22 +313,22 @@ def broadcast( def neighbors( offset_literal: it_ts.OffsetLiteralType, it: it_ts.IteratorType, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> ts.ListType: assert isinstance(offset_literal, it_ts.OffsetLiteralType) and isinstance( offset_literal.value, str ) assert isinstance(it, it_ts.IteratorType) conn_type = common.get_offset_type(offset_provider_type, offset_literal.value) - assert isinstance(conn_type, common.NeighborConnectivityType) - return ts.ListType(element_type=it.element_type, offset_type=conn_type.neighbor_dim) + assert isinstance(conn_type, common.NeighborTableType) + return ts.ListType(element_type=it.element_type, offset_type=conn_type.domain[1]) @_register_builtin_type_synthesizer def lift(stencil: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer def apply_lift( - *its: it_ts.IteratorType, offset_provider_type: common.OffsetProviderType + *its: it_ts.IteratorType, offset_provider_type: common.TableTypes ) -> it_ts.IteratorType: assert all(isinstance(it, it_ts.IteratorType) for it in its) stencil_args = [ @@ -406,7 +406,7 @@ def _canonicalize_nb_fields( Examples: >>> class Vertex(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> input_field = ts.FieldType( ... dims=[ ... Vertex, @@ -434,7 +434,7 @@ def _canonicalize_nb_fields( defined_dims = [] neighbor_dim = None for dim in input_dims: - if dim.kind == common.DimensionKind.LOCAL: + if common.is_local_dimension(dim): assert neighbor_dim is None neighbor_dim = dim else: @@ -451,7 +451,7 @@ def _canonicalize_nb_fields( def _resolve_dimensions( input_dims: list[common.Dimension], shift_tuple: tuple[itir.OffsetLiteral | itir.CartesianOffset, ...], - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> list[common.Dimension]: """ Resolves the final dimensions by applying shifts from the given shift tuple. @@ -489,21 +489,16 @@ def _resolve_dimensions( ... itir.OffsetLiteral(value="V2E"), ... itir.OffsetLiteral(value=0), ... ) + >>> def table_type(domain, codomain, max_neighbors): # of tables no declaration names + ... structure = common.ConnectivityType( + ... domain=domain, codomain=codomain, skip_value=None, dtype=None + ... ) + ... return common.NeighborTableType( + ... connectivity=structure, dtype=None, skip_value=None, max_neighbors=max_neighbors + ... ) >>> offset_provider_type = { - ... "C2V": common.NeighborConnectivityType( - ... domain=(Cell, C2V), - ... codomain=Vertex, - ... skip_value=None, - ... dtype=None, - ... max_neighbors=3, - ... ), - ... "V2E": common.NeighborConnectivityType( - ... domain=(Vertex, V2E), - ... codomain=Edge, - ... skip_value=None, - ... dtype=None, - ... max_neighbors=4, - ... ), + ... "C2V": table_type((Cell, C2V), Vertex, 3), + ... "V2E": table_type((Vertex, V2E), Edge, 4), ... } >>> _resolve_dimensions(input_dims, shift_tuple, offset_provider_type) [gt4py.next.iterator.type_system.type_synthesizer.Cell[horizontal], gt4py.next.iterator.type_system.type_synthesizer.K[vertical]] @@ -555,7 +550,7 @@ def _resolve_dimensions( off_literal.value, str ) offset_type = common.get_offset_type(offset_provider_type, off_literal.value) - if isinstance(offset_type, common.NeighborConnectivityType): + if isinstance(offset_type, common.NeighborTableType): if resolved_dim == offset_type.codomain: # Check if input fits to offset resolved_dim = offset_type.domain[0] # Update input_dim for next iteration else: @@ -571,7 +566,7 @@ def as_fieldop( stencil: TypeSynthesizer, domain: Optional[ts.DomainType] = None, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> TypeSynthesizer: @type_synthesizer def applied_as_fieldop( @@ -647,7 +642,7 @@ def scan( @type_synthesizer def apply_scan( - *its: it_ts.IteratorType, offset_provider_type: common.OffsetProviderType + *its: it_ts.IteratorType, offset_provider_type: common.TableTypes ) -> ts.DataType: result = scan_pass(init, *its, offset_provider_type=offset_provider_type) assert isinstance(result, ts.DataType) @@ -659,9 +654,7 @@ def apply_scan( @_register_builtin_type_synthesizer def map_list(op: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer - def applied_map( - *args: ts.ListType, offset_provider_type: common.OffsetProviderType - ) -> ts.ListType: + def applied_map(*args: ts.ListType, offset_provider_type: common.TableTypes) -> ts.ListType: assert len(args) > 0 assert all(isinstance(arg, ts.ListType) for arg in args) arg_el_types = [arg.element_type for arg in args] @@ -682,9 +675,7 @@ def _make_tuple_map_synthesizer( def tuple_map_synthesizer(op: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer - def applied_map( - arg: ts.TupleType, offset_provider_type: common.OffsetProviderType - ) -> ts.TupleType: + def applied_map(arg: ts.TupleType, offset_provider_type: common.TableTypes) -> ts.TupleType: if not isinstance(arg, ts.TupleType): raise TypeError( f"'{builtin_name}' requires a 'TupleType' argument, got '{type(arg).__name__}'." @@ -718,7 +709,7 @@ def applied_map( @_register_builtin_type_synthesizer def reduce(op: TypeSynthesizer, init: ts.TypeSpec) -> TypeSynthesizer: @type_synthesizer - def applied_reduce(*args: ts.ListType, offset_provider_type: common.OffsetProviderType): + def applied_reduce(*args: ts.ListType, offset_provider_type: common.TableTypes): assert all(isinstance(arg, ts.ListType) for arg in args) assert any( arg.offset_type is not None for arg in args @@ -731,7 +722,7 @@ def applied_reduce(*args: ts.ListType, offset_provider_type: common.OffsetProvid @_register_builtin_type_synthesizer -def shift(*offset_literals, offset_provider_type: common.OffsetProviderType) -> TypeSynthesizer: +def shift(*offset_literals, offset_provider_type: common.TableTypes) -> TypeSynthesizer: @type_synthesizer def apply_shift( it: it_ts.IteratorType | ts.DeferredType, @@ -754,7 +745,7 @@ def apply_shift( assert isinstance(offset_axis, it_ts.OffsetLiteralType) assert isinstance(offset_axis.value, str) type_ = common.get_offset_type(offset_provider_type, offset_axis.value) - assert isinstance(type_, common.NeighborConnectivityType) + assert isinstance(type_, common.NeighborTableType) source_dim, target_dim = type_.domain[0], type_.codomain found = False diff --git a/src/gt4py/next/otf/arguments.py b/src/gt4py/next/otf/arguments.py index 67c9f2bdc4..c0f86a96f6 100644 --- a/src/gt4py/next/otf/arguments.py +++ b/src/gt4py/next/otf/arguments.py @@ -135,7 +135,7 @@ class CompileTimeArgs: args: tuple[ts.TypeSpec, ...] kwargs: dict[str, ts.TypeSpec] - offset_provider: common.OffsetProvider # TODO(havogt): replace with common.OffsetProviderType once the temporary pass doesn't require the runtime information + offset_provider: common.OffsetProvider # TODO(havogt): replace with common.TableTypes once the temporary pass doesn't require the runtime information column_axis: Optional[common.Dimension] #: A mapping from an argument descriptor type to a context containing the actual descriptors. #: If an argument or element of an argument has no descriptor, the respective value is `None`. @@ -144,7 +144,7 @@ class CompileTimeArgs: argument_descriptor_contexts: ArgStaticDescriptorsContextsByType @property - def offset_provider_type(self) -> common.OffsetProviderType: + def offset_provider_type(self) -> common.TableTypes: return common.offset_provider_to_type(self.offset_provider) @classmethod diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index b71fdd9df9..f614b0956e 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -91,7 +91,7 @@ def compile_variant_hook( key: CompiledProgramsKey, backend: gtx_backend.Backend, argument_descriptors: ArgStaticDescriptorsByType, - offset_provider: common.OffsetProviderType | common.OffsetProvider, + offset_provider: common.TableTypes | common.OffsetProvider, ) -> None: """Callback hook invoked before compiling a program variant.""" @@ -627,7 +627,7 @@ def _finish_compilation_job(self, key: CompiledProgramsKey) -> bool: def _compile_variant( self, argument_descriptors: ArgStaticDescriptorsByType, - offset_provider: common.OffsetProviderType | common.OffsetProvider, + offset_provider: common.TableTypes | common.OffsetProvider, #: tuple consisting of the types of the positional and keyword arguments. arg_specialization_info: tuple[tuple[ts.TypeSpec, ...], dict[str, ts.TypeSpec]] | None = None, @@ -636,9 +636,9 @@ def _compile_variant( call_key: CompiledProgramsKey | None = None, ) -> None: if not common.is_offset_provider(offset_provider): - if common.is_offset_provider_type(offset_provider): + if common.is_table_types(offset_provider): raise ValueError( - "Variant compilation of programs with 'OffsetProviderType' is not yet supported." + "Variant compilation of programs with 'TableTypes' is not yet supported." ) else: raise ValueError(f"Invalid 'offset_provider': {offset_provider}") @@ -709,11 +709,11 @@ def _compile_variant( # domains and of scans. def compile( self, - offset_providers: list[common.OffsetProvider | common.OffsetProviderType], + offset_providers: list[common.OffsetProvider | common.TableTypes], **static_args: list[ScalarOrTupleOfScalars], ) -> None: """ - Compiles the program for all combinations of static arguments and the given 'OffsetProviderType'. + Compiles the program for all combinations of static arguments and the given 'TableTypes'. Note: In case you want to compile for specific combinations of static arguments (instead of the combinatoral), you can call compile multiples times. diff --git a/src/gt4py/next/otf/options.py b/src/gt4py/next/otf/options.py index 4f77d44586..6f0d7ac7a7 100644 --- a/src/gt4py/next/otf/options.py +++ b/src/gt4py/next/otf/options.py @@ -32,7 +32,7 @@ class CompilationOptions: #: when jitting is enabled, or on a call to `compile`. static_params: Sequence[str] | None = None - # TODO(ricoh): replace with common.OffsetProviderType once the temporary pass doesn't require the runtime information + # TODO(ricoh): replace with common.TableTypes once the temporary pass doesn't require the runtime information #: A dictionary holding static/compile-time information about the offset providers. #: For now, it is used for ahead of time compilation in DaCe orchestrated programs, #: i.e. DaCe programs that call GT4Py Programs -SDFGConvertible interface-. diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index 68591309c6..3c4f2ce1f2 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -68,7 +68,7 @@ def _process_regular_arguments( self, program: itir.Program, arg_types: tuple[ts.TypeSpec, ...], - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> tuple[list[interface.Parameter], list[str]]: parameters: list[interface.Parameter] = [] arg_exprs: list[str] = [] @@ -86,27 +86,32 @@ def _process_regular_arguments( isinstance( dim, fbuiltins.FieldOffset ) # TODO(havogt): remove support for FieldOffset as Dimension - or dim.kind is common.DimensionKind.LOCAL + or common.is_local_dimension(dim) ): # translate sparse dimensions to tuple dtype # NOTE: the tag is the offset-provider key, and its mangled form names the # `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) - assert isinstance(connectivity, common.NeighborConnectivityType) + 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.NeighborTableType) size = connectivity.max_neighbors arg = f"gridtools::sid::dimension_to_tuple_like({arg})" arg_exprs.append(arg) return parameters, arg_exprs def _process_connectivity_args( - self, offset_provider_type: common.OffsetProviderType + self, offset_provider_type: common.TableTypes ) -> tuple[list[interface.Parameter], list[str]]: parameters: list[interface.Parameter] = [] arg_exprs: list[str] = [] for name, connectivity_type in offset_provider_type.items(): - if isinstance(connectivity_type, common.NeighborConnectivityType): + if isinstance(connectivity_type, common.NeighborTableType): if connectivity_type.dtype.scalar_type not in [np.int32, np.int64]: raise ValueError( "Neighbor table indices must be of type 'np.int32' or 'np.int64'." @@ -142,7 +147,7 @@ def _process_connectivity_args( ) else: raise AssertionError( - f"Expected offset provider type '{name}' to be a 'NeighborConnectivityType', " + f"Expected offset provider type '{name}' to be a 'NeighborTableType', " f"got '{type(connectivity_type).__name__}'." ) @@ -151,7 +156,7 @@ def _process_connectivity_args( def _preprocess_program( self, program: itir.Program, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, ) -> itir.Program: return pass_manager.apply_common_transforms( program, @@ -165,7 +170,7 @@ def _preprocess_program( def generate_stencil_source( self, program: itir.Program, - offset_provider: common.OffsetProvider | common.OffsetProviderType, + offset_provider: common.OffsetProvider | common.TableTypes, column_axis: Optional[common.Dimension], ) -> str: if self.enable_itir_transforms: diff --git a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py index 38fe8fa1c6..c893be66e9 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py @@ -159,7 +159,7 @@ def _collect_dimensions_from_params( def _collect_offset_definitions( node: itir.Node, grid_type: common.GridType, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, ) -> dict[str, TagDefinition]: offset_definitions = {} offset_provider_type = {**offset_provider_type} @@ -188,17 +188,17 @@ def _collect_offset_definitions( ) for offset_name, connectivity_type in offset_provider_type.items(): - if isinstance(connectivity_type, common.NeighborConnectivityType): + if isinstance(connectivity_type, common.NeighborTableType): assert grid_type == common.GridType.UNSTRUCTURED offset_definitions[offset_name] = TagDefinition( name=Sym(id=common.codegen_name(offset_name)) ) - if offset_name != connectivity_type.neighbor_dim.tag: - offset_definitions[connectivity_type.neighbor_dim.tag] = TagDefinition( - name=Sym(id=common.codegen_name(connectivity_type.neighbor_dim.tag)) + if offset_name != connectivity_type.domain[1].tag: + offset_definitions[connectivity_type.domain[1].tag] = TagDefinition( + name=Sym(id=common.codegen_name(connectivity_type.domain[1].tag)) ) - for dim in [connectivity_type.source_dim, connectivity_type.codomain]: + for dim in [connectivity_type.domain[0], connectivity_type.codomain]: if dim.kind != common.DimensionKind.HORIZONTAL: raise NotImplementedError() offset_definitions[dim.tag] = TagDefinition( @@ -206,7 +206,7 @@ def _collect_offset_definitions( ) else: raise AssertionError( - "Elements of the offset provider type need to be a 'NeighborConnectivityType'." + "Elements of the offset provider type need to be a 'NeighborTableType'." ) return offset_definitions @@ -339,7 +339,7 @@ class GTFN_lowering(eve.NodeTranslator, eve.VisitorWithSymbolTableTrait): } _unary_op_map: ClassVar[dict[str, str]] = {"not_": "!"} - offset_provider_type: common.OffsetProviderType + offset_provider_type: common.TableTypes column_axis: Optional[common.Dimension] grid_type: common.GridType @@ -354,7 +354,7 @@ def apply( cls, node: itir.Program, *, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, column_axis: Optional[common.Dimension], ) -> Program: if not isinstance(node, itir.Program): @@ -506,7 +506,7 @@ def _visit_unstructured_domain(self, node: itir.FunCall, **kwargs: Any) -> Node: for o in shift_offsets: if o in self.offset_provider_type and isinstance( common.get_offset_type(self.offset_provider_type, o), - common.NeighborConnectivityType, + common.NeighborTableType, ): # `o` is an offset-provider key, i.e. a qualified tag: mangle it exactly as # its `TagDefinition` was, or the reference names an undeclared tag type. diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index c4b1526bdc..f195670437 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -87,7 +87,12 @@ class DataflowBuilder(Protocol): """Visitor interface to build a dataflow subgraph.""" @abc.abstractmethod - def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: ... + def get_offset_provider_type(self, offset: str) -> gtx_common.NeighborTableType: ... + + @abc.abstractmethod + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + """The offset of a connectivity over `local_dim`, see `common.connectivity_key_over`.""" + ... @abc.abstractmethod def unique_nsdfg_name(self, prefix: str) -> str: ... @@ -555,21 +560,24 @@ class GTIRToSDFG(eve.NodeVisitor, SDFGBuilder): from where to continue building the SDFG. """ - offset_provider_type: gtx_common.OffsetProviderType + offset_provider_type: gtx_common.TableTypes column_axis: Optional[gtx_common.Dimension] uids: gtx_utils.IDGeneratorPool = dataclasses.field( init=False, repr=False, default_factory=lambda: gtx_utils.IDGeneratorPool() ) - def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: + def get_offset_provider_type(self, offset: str) -> gtx_common.NeighborTableType: return gtx_common.get_offset_type(self.offset_provider_type, offset) + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + return gtx_common.connectivity_key_over(self.offset_provider_type, local_dim) + def make_field( self, data_node: dace_nodes.AccessNode, data_type: ts.FieldType, ) -> gtir_to_sdfg_types.FieldopData: - local_dims = [dim for dim in data_type.dims if dim.kind == gtx_common.DimensionKind.LOCAL] + local_dims = [dim for dim in data_type.dims if gtx_common.is_local_dimension(dim)] if len(local_dims) == 0: # do nothing: the field domain consists of all global dimensions field_type = data_type @@ -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 @@ -837,9 +847,9 @@ def _make_array_shape_and_strides( neighbor_table_types = gtx_dace_args.filter_connectivity_types(self.offset_provider_type) shape = [] for dim in dims: - if dim.kind == gtx_common.DimensionKind.LOCAL: + if gtx_common.is_local_dimension(dim): # for local dimension, the size is taken from the associated connectivity type - shape.append(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)) @@ -925,7 +935,7 @@ def _add_storage( all_dims = gt_type.dims else: # for 'ts.ListType' use 'offset_type' as local dimension assert gt_type.dtype.offset_type is not None - assert gt_type.dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(gt_type.dtype.offset_type) assert isinstance(gt_type.dtype.element_type, ts.ScalarType) dc_dtype = gtx_dace_args.as_dace_type(gt_type.dtype.element_type) all_dims = gtx_common.order_dimensions([*gt_type.dims, gt_type.dtype.offset_type]) @@ -1030,7 +1040,7 @@ def _add_sdfg_params( self.offset_provider_type ).items(): gt_type = ts.FieldType( - dims=[connectivity_type.source_dim, connectivity_type.neighbor_dim], + dims=[connectivity_type.domain[0], connectivity_type.domain[1]], dtype=tt.from_dtype(connectivity_type.dtype), ) # We store all connectivity tables as transient arrays here; later, while building @@ -1386,7 +1396,7 @@ def visit_SymRef( def lower_program_to_sdfg( ir: gtir.Program, - offset_provider_type: gtx_common.OffsetProviderType, + offset_provider_type: gtx_common.TableTypes, column_axis: Optional[gtx_common.Dimension] = None, ) -> dace.SDFG: """ diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py index 737648b911..0eb550936c 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py @@ -254,8 +254,10 @@ 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) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(local_dim) + ) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) local_idx = gtx_common.order_dimensions([*output_dims, local_dim]).index(local_dim) output_shape.insert(local_idx, offset_provider_type.max_neighbors) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index a28aad41c3..76cbfeb05b 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -80,7 +80,7 @@ class ValueExpr: def __post_init__(self) -> None: if isinstance(self.gt_dtype, ts.ListType): assert self.gt_dtype.offset_type is not None - assert self.gt_dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(self.gt_dtype.offset_type) @dataclasses.dataclass(frozen=True) @@ -106,7 +106,7 @@ def gt_dtype(self) -> ts.ScalarType | ts.ListType: def __post_init__(self) -> None: if isinstance(self.gt_dtype, ts.ListType): assert self.gt_dtype.offset_type is not None - assert self.gt_dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(self.gt_dtype.offset_type) @dataclasses.dataclass(frozen=True) @@ -146,7 +146,7 @@ def __post_init__(self) -> None: gtx_common.check_dims([dim for dim, _ in self.field_domain]) if isinstance(self.gt_dtype, ts.ListType): assert self.gt_dtype.offset_type is not None - assert self.gt_dtype.offset_type.kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(self.gt_dtype.offset_type) assert all(dim != self.gt_dtype.offset_type for dim, _ in self.field_domain) def get_field_type(self) -> ts.FieldType: @@ -762,12 +762,14 @@ 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), - gtx_common.NeighborConnectivityType, + self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(local_dim) + ), + gtx_common.NeighborTableType, ) # find position of the local dimension in the field layout assert isinstance(arg_desc, dace.data.Array) - assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in field_dims) + assert all(not gtx_common.is_local_dimension(dim) for dim in field_dims) extended_dims = gtx_common.order_dimensions([*field_dims, local_dim]) local_dim_pos = extended_dims.index(local_dim) inner_desc = dace.data.Array( @@ -1077,7 +1079,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: offset = node.args[0].value assert isinstance(offset, str) conn_type = self.subgraph_builder.get_offset_provider_type(offset) - assert isinstance(conn_type, gtx_common.NeighborConnectivityType) + assert isinstance(conn_type, gtx_common.NeighborTableType) it = self.visit(node.args[1]) assert isinstance(it, IteratorExpr) @@ -1090,8 +1092,8 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: origin for dim, origin in it.field_domain if dim == conn_type.codomain ) # make sure that the iterator can access the connectivity table - assert conn_type.source_dim in it.indices - conn_source_index = it.indices[conn_type.source_dim] + assert conn_type.domain[0] in it.indices + conn_source_index = it.indices[conn_type.domain[0]] assert isinstance(conn_source_index, SymbolExpr) # initially, the storage for the connectivty tables is created as transient; @@ -1128,7 +1130,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: ) # The layout of connectivity tables is known. assert len(conn_type.domain) == 2 - assert conn_type.domain[1].kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(conn_type.domain[1]) conn_slice = self._construct_local_view( MemletExpr( dc_node=self.state.add_access(conn_data), @@ -1153,7 +1155,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: # tag. The latter named a local dimension after the *offset*, which only coincided with # the real one under the old `V2EDim = Dimension("V2E")` convention, and under nominal # identity (ADR 0029) a tag string cannot be turned back into a dimension at all. - offset_type = conn_type.neighbor_dim + offset_type = conn_type.domain[1] neighbor_idx = gtir_to_sdfg_utils.get_map_variable(offset_type) index_connector = "__index" @@ -1307,7 +1309,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: tasklet_expression = f"{output_connector} = {fun_python_code}" input_args = [self.visit(arg) for arg in node.args] - input_conn_types: dict[gtx_common.Dimension, gtx_common.NeighborConnectivityType] = {} + input_conn_types: dict[gtx_common.Dimension, gtx_common.NeighborTableType] = {} for input_arg in input_args: assert isinstance(input_arg.gt_dtype, ts.ListType) assert input_arg.gt_dtype.offset_type is not None @@ -1315,8 +1317,10 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if offset_type is gtx_common.ConstList: # this input argument is the result of `make_const_list` continue - offset_provider_t = self.subgraph_builder.get_offset_provider_type(offset_type.tag) - assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) + 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.NeighborTableType) input_conn_types[offset_type] = offset_provider_t if len(input_conn_types) == 0: @@ -1369,15 +1373,17 @@ 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 - origin_map_index = gtir_to_sdfg_utils.get_map_variable(conn_type.source_dim) + origin_map_index = gtir_to_sdfg_utils.get_map_variable(conn_type.domain[0]) # The layout of connectivity tables is known. assert len(conn_type.domain) == 2 - assert conn_type.domain[1].kind == gtx_common.DimensionKind.LOCAL + assert gtx_common.is_local_dimension(conn_type.domain[1]) conn_slice = self._construct_local_view( MemletExpr( dc_node=self.state.add_access(conn_data), @@ -1432,9 +1438,9 @@ 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) + assert isinstance(offset_provider_t, gtx_common.NeighborTableType) local_size = offset_provider_t.max_neighbors map_index = gtir_to_sdfg_utils.get_map_variable(list_type.offset_type) @@ -1472,8 +1478,10 @@ 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) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + 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.NeighborTableType) inp_conn = "_in" outp_conn = "_out" @@ -1484,7 +1492,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( @@ -1499,7 +1509,7 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: ) self.state.add_node(reduce_node) - origin_map_index = gtir_to_sdfg_utils.get_map_variable(offset_provider_type.source_dim) + origin_map_index = gtir_to_sdfg_utils.get_map_variable(offset_provider_type.domain[0]) self._add_input_data_edge( self.state.add_access(connectivity), dace_subsets.Range.from_string( @@ -1689,7 +1699,7 @@ def _make_dynamic_neighbor_offset( def _make_unstructured_shift( self, it: IteratorExpr, - conn_type: gtx_common.NeighborConnectivityType, + conn_type: gtx_common.NeighborTableType, conn_node: dace_nodes.AccessNode, offset_expr: DataExpr, ) -> IteratorExpr: @@ -1697,19 +1707,19 @@ def _make_unstructured_shift( # make sure that the field can be dereferenced with the given connectivity type assert any(dim == conn_type.codomain for dim, _ in it.field_domain) # make sure that the iterator can access the connectivity table - assert conn_type.source_dim in it.indices - conn_source_index = it.indices[conn_type.source_dim] + assert conn_type.domain[0] in it.indices + conn_source_index = it.indices[conn_type.domain[0]] assert isinstance(conn_source_index, SymbolExpr) shifted_indices = { - dim: idx for dim, idx in it.indices.items() if dim != conn_type.source_dim + dim: idx for dim, idx in it.indices.items() if dim != conn_type.domain[0] } if isinstance(offset_expr, SymbolExpr): # use memlet to retrieve the neighbor index shifted_indices[conn_type.codomain] = MemletExpr( dc_node=conn_node, gt_field=ts.FieldType( - dims=[conn_type.source_dim], + dims=[conn_type.domain[0]], dtype=ts.ListType( element_type=tt.from_dtype(conn_type.dtype), offset_type=gtx_common.ConstList, @@ -1751,7 +1761,7 @@ def _visit_shift(self, node: gtir.FunCall) -> IteratorExpr: offset_provider_type = self.subgraph_builder.get_offset_provider_type( offset_provider_arg.value ) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) # a named offset → unstructured shift; the offset value may be a static # `OffsetLiteral` or a dynamic offset (handled by `_make_unstructured_shift`). # initially, the storage for the connectivity tables is created as transient; diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index e265377a45..d779e71fef 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -134,7 +134,7 @@ def _create_field_operator_impl( assert isinstance(dataflow_output_desc, dace.data.Array) assert len(dataflow_output_desc.shape) == 1 # extend the array with the local dimensions added by the field operator (e.g. `neighbors`) - assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in field_dims) + assert all(not gtx_common.is_local_dimension(dim) for dim in field_dims) assert output_edge.result.gt_dtype.offset_type is not None local_dim = output_edge.result.gt_dtype.offset_type # construct the full subset according to the canonical field domain @@ -325,8 +325,10 @@ 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) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + 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.NeighborTableType) shape = [*shape, offset_provider_type.max_neighbors] out, _ = sdfg_builder.add_temp_array(ctx.sdfg, shape, dtype) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index 18b48987cb..911078c9c7 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -384,8 +384,10 @@ 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) - assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(offset_type) + ) + assert isinstance(offset_provider_type, gtx_common.NeighborTableType) list_size = offset_provider_type.max_neighbors return [scan_column_size, dace.symbolic.SymExpr(list_size)] diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py index 47baf902ec..037c2989bc 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_types.py @@ -73,7 +73,7 @@ def get_local_view( # The invariant below is ensured by calling `make_field()` to construct `FieldopData`. # The `make_field` constructor converts any local dimension, if present, to `ListType` # element type, while leaving the field domain with all global dimensions. - assert all(dim.kind != gtx_common.DimensionKind.LOCAL for dim in self.gt_type.dims) + assert all(not gtx_common.is_local_dimension(dim) for dim in self.gt_type.dims) domain_dims = [domain_range.dim for domain_range in domain] domain_indices = gtir_domain.get_element_subset( domain_dims, origin=None diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py index a2117a87a2..06c0ec272a 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_utils.py @@ -50,8 +50,9 @@ def get_map_variable(dim: gtx_common.Dimension) -> str: # fusion and map splitting rely on the names of the map variables to match the field # dimensions and decide whether two maps have the same iteration space. dim = gtx_common.as_non_staggered(dim) - suffix = "dim" if dim.kind == gtx_common.DimensionKind.LOCAL else "" - return f"i_{gtx_common.codegen_name(dim.tag)}_gtx_{dim.kind}{suffix}" + # NOTE: a local dimension has no `kind`; it keeps the name it had when `LOCAL` was a kind. + kind = "localdim" if gtx_common.is_local_dimension(dim) else str(dim.kind) + return f"i_{gtx_common.codegen_name(dim.tag)}_gtx_{kind}" def make_tasklet_connector_for(name: str) -> str: diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py index 5b0b0601fc..55d3c77ada 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -60,7 +60,7 @@ def connectivity_identifier(name: str) -> str: def is_connectivity_identifier( - name: str, offset_provider_type: gtx_common.OffsetProviderType | None = None + name: str, offset_provider_type: gtx_common.TableTypes | None = None ) -> bool: if (m := CONNECTIVITY_INDENTIFIER_RE.match(name)) is None: return False @@ -76,7 +76,7 @@ def _field_symbol( field_name: str, dim: gtx_common.Dimension, sym: Literal["size", "stride"], - offset_provider_type: gtx_common.OffsetProviderType | None, + offset_provider_type: gtx_common.TableTypes | None, ) -> dace.symbol: if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is None: name = f"__{field_name}_{gtx_common.codegen_name(dim.tag)}_{sym}" @@ -85,10 +85,10 @@ def _field_symbol( offset = gtx_common.from_codegen_name(m[1]) assert offset in offset_provider_type conn_type = offset_provider_type[offset] - assert isinstance(conn_type, gtx_common.NeighborConnectivityType) - if dim == conn_type.source_dim: + assert isinstance(conn_type, gtx_common.NeighborTableType) + if dim == conn_type.domain[0]: name = f"__{field_name}_source_{sym}" - elif dim == conn_type.neighbor_dim: + elif dim == conn_type.domain[1]: name = f"__{field_name}_neighbor_{sym}" else: raise ValueError(f"Unexpect dimension '{dim}' for '{offset}' connectivity.") @@ -98,7 +98,7 @@ def _field_symbol( def field_size_symbol( field_name: str, dim: gtx_common.Dimension, - offset_provider_type: gtx_common.OffsetProviderType, + offset_provider_type: gtx_common.TableTypes, ) -> dace.symbol: return _field_symbol(field_name, dim, "size", offset_provider_type) @@ -106,11 +106,32 @@ def field_size_symbol( def field_stride_symbol( field_name: str, dim: gtx_common.Dimension, - offset_provider_type: gtx_common.OffsetProviderType | None = None, + offset_provider_type: gtx_common.TableTypes | None = None, ) -> dace.symbol: return _field_symbol(field_name, dim, "stride", offset_provider_type) +def local_dimension_size( + field_name: str, + dim: gtx_common.Dimension, + neighbor_table_types: dict[str, gtx_common.NeighborTableType], +) -> 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.domain[1] == 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) @@ -130,15 +151,15 @@ def range_stop_symbol(field_name: str, dim: gtx_common.Dimension) -> dace.symbol def filter_connectivity_types( - offset_provider_type: gtx_common.OffsetProviderType, -) -> dict[str, gtx_common.NeighborConnectivityType]: + offset_provider_type: gtx_common.TableTypes, +) -> dict[str, gtx_common.NeighborTableType]: """ - Filter offset provider types of type `NeighborConnectivityType`. + Filter offset provider types of type `NeighborTableType`. In other words, filter out the cartesian offset providers. """ return { offset: conn for offset, conn in offset_provider_type.items() - if isinstance(conn, gtx_common.NeighborConnectivityType) + if isinstance(conn, gtx_common.NeighborTableType) } diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index ee96fbeff2..df2ceb6670 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -31,7 +31,7 @@ def find_constant_symbols( ir: itir.Program, sdfg: dace.SDFG, - offset_provider_type: common.OffsetProviderType, + offset_provider_type: common.TableTypes, disable_field_origin_on_program_arguments: bool, unstructured_horizontal_has_unit_stride: bool, ) -> dict[str, int]: @@ -56,13 +56,13 @@ def find_constant_symbols( # Same for connectivity tables, for which the first dimension is always horizontal for offset, conn_type in offset_provider_type.items(): if ( - isinstance(conn_type, common.NeighborConnectivityType) + isinstance(conn_type, common.NeighborTableType) and (conn_id := gtx_dace_args.connectivity_identifier(offset)) in sdfg.arrays ): assert not sdfg.arrays[conn_id].transient - assert conn_type.source_dim.kind == common.DimensionKind.HORIZONTAL + assert conn_type.domain[0].kind == common.DimensionKind.HORIZONTAL sdfg_stride_symbol = gtx_dace_args.field_stride_symbol( - conn_id, conn_type.source_dim, offset_provider_type + conn_id, conn_type.domain[0], offset_provider_type ) constant_symbols[sdfg_stride_symbol.name] = 1 diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 453d5f8dad..056473db91 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -409,7 +409,7 @@ def is_local_field(type_: ts.FieldType) -> bool: Examples: >>> class V(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> is_local_field( ... ts.FieldType(dims=[V, V2E], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64)) ... ) @@ -417,7 +417,7 @@ def is_local_field(type_: ts.FieldType) -> bool: >>> is_local_field(ts.FieldType(dims=[V], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64))) False """ - return any(dim.kind == common.DimensionKind.LOCAL for dim in type_.dims) + return any(common.is_local_dimension(dim) for dim in type_.dims) def contains_local_field(type_: ts.TypeSpec) -> bool: @@ -533,7 +533,7 @@ def is_concretizable(symbol_type: ts.TypeSpec, to_type: ts.TypeSpec) -> bool: True >>> is_concretizable( - ... ts.DeferredType(constraint=ts.OffsetType), + ... ts.DeferredType(constraint=ts.ShiftType), ... to_type=ts.FieldType(dtype=ts.ScalarType(kind=ts.ScalarKind.BOOL), dims=[]), ... ) False @@ -587,7 +587,7 @@ def promote( >>> promoted.dims == [I, J, K] and promoted.dtype == dtype True - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> list_dtype = ts.ListType(element_type=dtype, offset_type=V2E) >>> promote( ... ts.FieldType(dims=[I], dtype=list_dtype), @@ -669,19 +669,19 @@ def return_type_field( except ValueError as ex: raise ValueError("Could not deduce return type of invalid remap operation.") from ex - if not isinstance(with_args[0], ts.OffsetType): - raise ValueError(f"First argument must be of type '{ts.OffsetType}', got '{with_args[0]}'.") + if not isinstance(with_args[0], ts.ShiftType): + raise ValueError(f"First argument must be of type '{ts.ShiftType}', got '{with_args[0]}'.") - source_dim = with_args[0].source - target_dims = with_args[0].target + codomain = with_args[0].codomain + domain_dims = with_args[0].domain new_dims = [] # TODO: This code does not handle ellipses for dimensions. Fix it. assert field_type.dims is not ... for d in field_type.dims: - if d != source_dim: + if d != codomain: new_dims.append(d) else: - new_dims.extend(target_dims) + new_dims.extend(domain_dims) return ts.FieldType(dims=new_dims, dtype=field_type.dtype) @@ -890,10 +890,10 @@ def function_signature_incompatibilities_field( yield f"Function takes at least 1 argument, but {len(args)} were given." return for arg in args: - if not isinstance(arg, ts.OffsetType): - yield f"Expected arguments to be of type '{ts.OffsetType}', got '{arg}'." + if not isinstance(arg, ts.ShiftType): + yield f"Expected arguments to be of type '{ts.ShiftType}', got '{arg}'." return - if len(args) > 1 and len(arg.target) > 1: + if len(args) > 1 and len(arg.domain) > 1: yield f"Function takes only 1 argument in unstructured case, but {len(args)} were given." return @@ -901,15 +901,15 @@ def function_signature_incompatibilities_field( yield f"Got unexpected keyword argument(s) '{', '.join(kwargs.keys())}'." return - source_dim = args[0].source # type: ignore[attr-defined] # ensured by loop above - target_dims = args[0].target # type: ignore[attr-defined] # ensured by loop above + codomain = args[0].codomain # type: ignore[attr-defined] # ensured by loop above + domain_dims = args[0].domain # type: ignore[attr-defined] # ensured by loop above assert field_type.dims is not ... - if field_type.dims and source_dim not in field_type.dims: + if field_type.dims and codomain not in field_type.dims: yield ( f"Incompatible offset can not shift field defined on " f"{', '.join([dim.__qualname__ for dim in field_type.dims])} from " - f"{source_dim.__qualname__} to target dim(s): " - f"{', '.join([dim.tag for dim in target_dims])}" + f"{codomain.__qualname__} to target dim(s): " + f"{', '.join([dim.tag for dim in domain_dims])}" ) diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 5ad8ea7105..648dfa866b 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -67,16 +67,28 @@ def __str__(self) -> str: return f"Index[{self.dim}]" -class OffsetType(TypeSpec): - # TODO(havogt): replace by ConnectivityType - source: common.Dimension - target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] +class ShiftType(TypeSpec): + """ + The type of a shift: it takes a field over `codomain` to a field over `domain`. + + `domain` has one dimension for a Cartesian shift (`KDim + 1`) and for a single neighbor + (`V2E[i]`), and two -- the connectivity's domain and its local dimension -- for all + neighbors (`V2E`). + """ + + codomain: common.Dimension + domain: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] #: The offset-provider key; `None` for the untagged Cartesian `Dim + offset`. tag: Optional[common.Tag] = None def __str__(self) -> str: tag = "" if self.tag is None else f"{self.tag}: " - return f"Offset[{tag}{self.source}, {self.target}]" + domain = ( + str(self.domain[0]) + if len(self.domain) == 1 + else f"({', '.join(str(dim) for dim in self.domain)})" + ) + return f"Shift[{tag}{self.codomain} -> {domain}]" class ScalarKind(eve_types.IntEnum): diff --git a/src/gt4py/next/type_system/type_translation.py b/src/gt4py/next/type_system/type_translation.py index 9722b27f5b..d66214c992 100644 --- a/src/gt4py/next/type_system/type_translation.py +++ b/src/gt4py/next/type_system/type_translation.py @@ -365,7 +365,7 @@ def from_value(value: Any) -> ts.TypeSpec: type_ = xtyping.infer_type(value, annotate_callable_kwargs=True) symbol_type = from_type_hint(type_) - if isinstance(symbol_type, (ts.DataType, ts.CallableType, ts.OffsetType, ts.DimensionType)): + if isinstance(symbol_type, (ts.DataType, ts.CallableType, ts.ShiftType, ts.DimensionType)): return symbol_type else: raise ValueError(f"Impossible to map '{value}' value to a 'Symbol'.") diff --git a/tests/next_tests/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/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 4b73f163a9..9fc0f7c1ae 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -184,7 +184,7 @@ class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) -class C2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2VDim(gtx.LocalDimensionIndex): ... V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) @@ -206,7 +206,7 @@ def sizes(self) -> tuple[int, int, int]: ... def offset_provider(self) -> common.OffsetProvider: ... @property - def offset_provider_type(self) -> common.OffsetProviderType: ... + def offset_provider_type(self) -> common.TableTypes: ... def simple_cartesian_grid( @@ -248,7 +248,7 @@ def num_edges(self) -> int: ... def offset_provider(self) -> common.OffsetProvider: ... @property - def offset_provider_type(self) -> common.OffsetProviderType: ... + def offset_provider_type(self) -> common.TableTypes: ... def simple_mesh(allocator) -> MeshDescriptor: diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py index e298515208..a510f383ce 100644 --- a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py +++ b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py @@ -42,7 +42,7 @@ class Cell(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(gtx.LocalDimensionIndex): ... C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) diff --git a/tests/next_tests/integration_tests/feature_tests/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..bce54f4e6a --- /dev/null +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py @@ -0,0 +1,199 @@ +# 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 dataclasses +import typing + +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): ... + + +#: A second connectivity over the same neighbor axis, bound to a different table. +class V2EShared(gtx.NeighborConnectivity[V, E]): + Local: typing.TypeAlias = V2E.Local + + +@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) + 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 + 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}, + 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, connectivity=V2E) -> np.ndarray: + return case.offset_provider[connectivity.offset_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]]) + + +@pytest.mark.uses_unstructured_shift +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 +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), + ) + + +@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), + ) + + +@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/feature_tests/instrumentation_tests/test_hooks.py b/tests/next_tests/integration_tests/feature_tests/instrumentation_tests/test_hooks.py index ce16667142..5e750f3c65 100644 --- a/tests/next_tests/integration_tests/feature_tests/instrumentation_tests/test_hooks.py +++ b/tests/next_tests/integration_tests/feature_tests/instrumentation_tests/test_hooks.py @@ -214,7 +214,7 @@ def custom_compile_variant_hook( key: gtx_typing.CompiledProgramsKey, backend: gtx_typing.Backend, argument_descriptors: dict[type, dict[str, Any]], - offset_provider: common.OffsetProviderType | common.OffsetProvider, + offset_provider: common.TableTypes | common.OffsetProvider, ) -> None: callback_results.append( ( diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py index bd6d613efc..e3e9a9286d 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py @@ -59,7 +59,7 @@ class Node(gtx.DimensionIndex): ... -class NeighDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class NeighDim(gtx.LocalDimensionIndex): ... def array_maker(*lists): diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py index bd675b5f51..afb2925f75 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py @@ -10,6 +10,7 @@ import pytest import gt4py.next as gtx +from gt4py.next import common from gt4py.next.iterator.builtins import deref, named_range, shift, unstructured_domain, as_fieldop from gt4py.next.iterator.runtime import set_at, fendef, fundef, offset @@ -17,7 +18,7 @@ from gt4py.next.iterator.embedded import StridedConnectivityField -class Dummy(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class Dummy(gtx.LocalDimensionIndex): ... class LocA(gtx.DimensionIndex): ... @@ -53,7 +54,10 @@ def test_strided_offset_provider(program_processor): program_processor, validate = program_processor LocA_size = 2 - max_neighbors = LocA2LocAB_offset_provider.__gt_type__().max_neighbors + # the table's type as bound under its key: a table's own `__gt_type__()` is only structural + max_neighbors = common.offset_provider_to_type({"O": LocA2LocAB_offset_provider})[ + "O" + ].max_neighbors LocAB_size = LocA_size * max_neighbors rng = np.random.default_rng() diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py index fdf8cf5114..2490d5ab91 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py @@ -23,7 +23,7 @@ reduce, shift, ) -from gt4py.next.iterator.runtime import fundef +from gt4py.next.iterator.runtime import fundef, offset from gt4py.next.program_processors.runners import gtfn from next_tests.toy_connectivity import ( @@ -434,3 +434,51 @@ def test_sparse_shifted_stencil_reduce(program_processor): if validate: assert np.allclose(out.asnumpy(), ref) + + +#: A second connectivity over `V2E`'s local dimension, under its own name; bound to the table +#: with its columns reversed. +V2E_SHARED = offset("V2EShared") + + +v2e_shared_arr = np.ascontiguousarray(v2e_arr[:, ::-1]) +v2e_shared_conn = gtx.as_connectivity( + domain={Vertex: v2e_shared_arr.shape[0], V2EDim: v2e_shared_arr.shape[1]}, + codomain=Edge, + data=v2e_shared_arr, +) + + +@fundef +def shift_through_sharer(in_edges): + return deref(shift(V2E_SHARED, 1)(in_edges)) + + +@fundef +def owner_times_sharer(in_edges): + return reduce(plus, 0)( + map_list(multiplies)(neighbors(V2E_SHARED, in_edges), neighbors(V2E, in_edges)) + ) + + +@pytest.mark.parametrize( + "stencil, ref", + [ + (shift_through_sharer, v2e_shared_arr[:, 1]), + (owner_times_sharer, np.sum(v2e_shared_arr * v2e_arr, axis=1)), + ], +) +def test_connectivity_sharing_a_local_dimension(program_processor, stencil, ref): + program_processor, validate = program_processor + inp = edge_index_field() + out = gtx.as_field([Vertex], np.zeros([9], dtype=inp.dtype)) + + run_processor( + stencil[{Vertex: range(0, 9)}], + program_processor, + inp, + out=out, + offset_provider={V2EDim.tag: v2e_conn, V2E_SHARED.value: v2e_shared_conn}, + ) + if validate: + assert np.allclose(out.asnumpy(), ref) diff --git a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py index 689d0d5f71..b7c81d55fd 100644 --- a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py +++ b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py @@ -43,29 +43,21 @@ 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)) -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,16 +130,14 @@ 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. - If `NeighborConnectivityType.neighbor_dim` did not match the `FieldOffset` value, + If the local dimension of the `NeighborTableType` did not match the `FieldOffset` value, gtfn would silently ignore the neighbor index, see https://github.com/GridTools/gridtools/pull/1814. """ @@ -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/toy_connectivity.py b/tests/next_tests/toy_connectivity.py index c368b04245..531ee0867f 100644 --- a/tests/next_tests/toy_connectivity.py +++ b/tests/next_tests/toy_connectivity.py @@ -21,16 +21,16 @@ class Edge(gtx.DimensionIndex): ... class Cell(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2VDim(gtx.LocalDimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(gtx.LocalDimensionIndex): ... -class V2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2VDim(gtx.LocalDimensionIndex): ... V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) diff --git a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py index 6a47073c8a..1cbb1818c2 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py @@ -20,6 +20,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, Domain, Field, @@ -50,10 +51,10 @@ class V(DimensionIndex): ... class E(DimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class V2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2V(LocalDimensionIndex): ... class C(DimensionIndex): ... @@ -62,7 +63,7 @@ class C(DimensionIndex): ... class K(CartesianAxisIndex): ... -class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E2CO(LocalDimensionIndex): ... class A(DimensionIndex): ... @@ -77,7 +78,7 @@ class X(CartesianAxisIndex): ... class Y(CartesianAxisIndex): ... -class L(DimensionIndex, kind=DimensionKind.LOCAL): ... +class L(LocalDimensionIndex): ... class S(DimensionIndex): ... @@ -98,7 +99,7 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... class C2V(DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py index 2817dd37dc..7c8c026f58 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py @@ -21,7 +21,7 @@ class VDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... class Dim(gtx.DimensionIndex): ... -class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class LocalDim(gtx.LocalDimensionIndex): ... CartesianOffset = gtx.FieldOffset("CartesianOffset", source=Dim, target=(Dim,)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index 2b8a8a4c4f..37c650d998 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py @@ -49,7 +49,7 @@ class Edge(gtx.DimensionIndex): ... class Vertex(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) @@ -63,7 +63,7 @@ class TDim(gtx.CartesianAxisIndex): ... #: An offset whose tag differs from the name of the Python variable it is bound to, and #: from the name of its local dimension. Lowering must emit the *tag*. -class RenamedV2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class RenamedV2EDim(gtx.LocalDimensionIndex): ... renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py index 290d2914fc..2856c31149 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py @@ -19,7 +19,7 @@ """ import gt4py.next as gtx -from gt4py.next import Dims, Dimension, DimensionIndex, float64, neighbor_sum +from gt4py.next import Dims, Dimension, DimensionIndex, LocalDimensionIndex, float64, neighbor_sum from gt4py.next.ffront import source_utils from gt4py.next.ffront.source_utils import get_closure_vars_from_function @@ -30,7 +30,7 @@ class Cell(DimensionIndex): ... class Edge(DimensionIndex): ... -class C2EDim(DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py index ca5cfd4c79..afa3bba76e 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py @@ -18,6 +18,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, Field, FieldOffset, @@ -51,7 +52,7 @@ class X(CartesianAxisIndex): ... class Y(CartesianAxisIndex): ... -class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class Y2XDim(LocalDimensionIndex): ... class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... @@ -72,7 +73,7 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... class IDim(CartesianAxisIndex): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py index 3c7157166d..aca610d25d 100644 --- a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py +++ b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py @@ -33,13 +33,13 @@ class Vertex(common.DimensionIndex): ... class Edge(common.DimensionIndex): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... -class V2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2VDim(common.LocalDimensionIndex): ... a_range = domain_utils.SymbolicRange(0, 10) @@ -315,7 +315,7 @@ def test_unstructured_translate(shift_chain, expected_end_domain): def test_unstructured_translate_with_symbolic_domain_sizes(as_type): # With `symbolic_domain_sizes` the translated range is taken from the provided size # expression instead of the connectivity table. This makes `translate` work for a type-only - # `OffsetProviderType` (which has no table) as well as a runtime `OffsetProvider`. + # `TableTypes` (which has no table) as well as a runtime `OffsetProvider`. offset_provider = { V2EDim.tag: constructors.as_connectivity( domain={Vertex: (0, 4), V2EDim: 1}, diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py index 869372ab78..33f6f71f43 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py @@ -30,7 +30,7 @@ class E(gtx.DimensionIndex): ... class V(gtx.DimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2VDim(gtx.LocalDimensionIndex): ... E2V = gtx.FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py index 7119c5fcb7..4be03cafb0 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py @@ -267,7 +267,7 @@ class IDim(gtx.CartesianAxisIndex): ... class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... -class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class LocalDim(gtx.LocalDimensionIndex): ... @pytest.mark.parametrize("dim, suffix", [(IDim, "ₕ"), (KDim, "ᵥ"), (LocalDim, "ₗ")]) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py b/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py index a742561830..d933548042 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py @@ -29,11 +29,16 @@ def foo(inp): return deref(inp) -connectivity = common.ConnectivityType( - domain=[dummy_origin, dummy_neighbor], - codomain=dummy_codomain, +connectivity = common.NeighborTableType( + connectivity=common.ConnectivityType( + domain=(dummy_origin, dummy_neighbor), + codomain=dummy_codomain, + skip_value=common._DEFAULT_SKIP_VALUE, + dtype=None, + ), skip_value=common._DEFAULT_SKIP_VALUE, dtype=None, + max_neighbors=1, ) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py index ff6890f906..abe6d9d28c 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py @@ -45,7 +45,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... float_i_field = ts.FieldType(dims=[IDim], dtype=float_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py index 250d3eabc3..5577e016be 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py @@ -17,7 +17,7 @@ from gt4py.next.type_system import type_specifications as ts -class Neighbor(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Neighbor(common.LocalDimensionIndex): ... class IDim(common.CartesianAxisIndex): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py index e4d1746ccc..ef95683b21 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py @@ -25,7 +25,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... class K(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py index 1028fd987f..ae5c1186fe 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py @@ -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 @@ -27,17 +27,22 @@ class dummy_neighbor(common.DimensionIndex): ... #: The local dimensions of the neighbor lists under test. Each one's `tag` is also its IR offset #: string and its offset-provider key: `UnrollReduce` looks a connectivity up by the local #: dimension of the list it reduces, so those three names must be a single string (ADR 0029). -class Dim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim(common.LocalDimensionIndex): ... -class Dim2(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim2(common.LocalDimensionIndex): ... def dummy_connectivity_type(max_neighbors: int, has_skip_values: bool): - return common.NeighborConnectivityType( - domain=[dummy_origin, dummy_neighbor], - codomain=dummy_codomain, - skip_value=common._DEFAULT_SKIP_VALUE if has_skip_values else None, + skip_value = common._DEFAULT_SKIP_VALUE if has_skip_values else None + return common.NeighborTableType( + connectivity=common.ConnectivityType( + domain=(dummy_origin, dummy_neighbor), + codomain=dummy_codomain, + skip_value=skip_value, + dtype=None, + ), + skip_value=skip_value, dtype=None, max_neighbors=max_neighbors, ) @@ -98,10 +103,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): diff --git a/tests/next_tests/unit_tests/otf_tests/test_runners.py b/tests/next_tests/unit_tests/otf_tests/test_runners.py index 0a0941d792..8974d7bad1 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_runners.py +++ b/tests/next_tests/unit_tests/otf_tests/test_runners.py @@ -31,7 +31,7 @@ class Vertex(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... @pytest.fixture diff --git a/tests/next_tests/unit_tests/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..96e121ac4f 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py @@ -18,3 +18,36 @@ def test_safe_replace_symbolic(): assert gtir_to_sdfg_utils.safe_replace_symbolic( dace.symbolic.pystr_to_symbolic("x*x + y"), symbol_mapping={"x": "y", "y": "x"} ) == dace.symbolic.pystr_to_symbolic("y*y + x") + + +def test_local_dimension_size(): + import numpy as np + + from gt4py._core import definitions as core_defs + from gt4py.next import common + from gt4py.next.program_processors.runners.dace import sdfg_args + + from next_tests.toy_connectivity import V2EDim, Vertex, Edge + + def conn_type(max_neighbors: int) -> common.NeighborTableType: + dtype = core_defs.dtype(np.int32) + return common.NeighborTableType( + connectivity=common.ConnectivityType( + domain=(Vertex, V2EDim), codomain=Edge, skip_value=None, dtype=dtype + ), + skip_value=None, + dtype=dtype, + max_neighbors=max_neighbors, + ) + + sharer_tag = "some.module.V2EShared" + table_types = {V2EDim.tag: conn_type(4), sharer_tag: conn_type(4)} + # a field finds the size in the table keyed by the local dimension + assert sdfg_args.local_dimension_size("a_field", V2EDim, table_types) == 4 + # a connectivity array has its own + conn_array = sdfg_args.connectivity_identifier(sharer_tag) + assert sdfg_args.local_dimension_size(conn_array, V2EDim, {sharer_tag: conn_type(4)}) == 4 + # a field over a local dimension bound only through a sharing connectivity + assert sdfg_args.local_dimension_size("a_field", V2EDim, {sharer_tag: conn_type(4)}) == 4 + with pytest.raises(KeyError, match="No connectivity over the local dimension"): + sdfg_args.local_dimension_size("a_field", V2EDim, {}) diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index c5daa24c9b..f307f9d5b3 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -22,6 +22,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, Domain, Infinity, @@ -58,19 +59,19 @@ class I(common.CartesianAxisIndex): ... class I_half(common.CartesianAxisIndex): ... -class C2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E(LocalDimensionIndex): ... -class V2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(LocalDimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C(LocalDimensionIndex): ... -class E2C2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2V(LocalDimensionIndex): ... class ECDim(DimensionIndex): ... @@ -634,7 +635,7 @@ def dimension_promotion_cases() -> list[ ( [[JDim, V2E], [IDim, E2C2V, KDim]], None, - "There are more than one dimension with DimensionKind 'LOCAL'.", + "There is more than one local dimension.", ), ([[JDim, V2E], [IDim, KDim]], [IDim, JDim, V2E, KDim], None), # a dimension and its staggered counterpart must not be promoted into the same field diff --git a/tests/next_tests/unit_tests/test_custom_layout_allocators.py b/tests/next_tests/unit_tests/test_custom_layout_allocators.py index 47c657c3d3..36ff69a3fe 100644 --- a/tests/next_tests/unit_tests/test_custom_layout_allocators.py +++ b/tests/next_tests/unit_tests/test_custom_layout_allocators.py @@ -29,13 +29,13 @@ class D2(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... class D0_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... -class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D1_local(common.LocalDimensionIndex): ... class D2_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... -class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D2_local(common.LocalDimensionIndex): ... class D1_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py new file mode 100644 index 0000000000..8e47de451f --- /dev/null +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -0,0 +1,675 @@ +# 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 ( + CartesianAxisIndex, + DimensionIndex, + DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, + NeighborTableType, +) +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(CartesianAxisIndex, 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__, + "typing": typing, + "DimensionIndex": DimensionIndex, + "DimensionKind": DimensionKind, + "LocalDimensionIndex": LocalDimensionIndex, + "NeighborConnectivity": NeighborConnectivity, + "Vertex": Vertex, + "Edge": Edge, + "KDim": KDim, + "V2E": V2E, + "ConstList": common.ConstList, + } + exec(textwrap.dedent(source), namespace) + return namespace + + +class TestDeclaration: + def test_owner_and_dimensions(self): + assert V2E.Local.owner is V2E + assert V2E.domain is Vertex + assert V2E.codomain is Edge + assert common.is_local_dimension(V2E.Local) + assert V2E.Local.kind is None + 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 common.is_local_dimension(LsqCoeff) + assert LsqCoeff.kind is None + + 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): ... + """, + "must declare its local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=5): + Local: typing.TypeAlias = V2E.Local + """, + "contradicts the local dimension it shares with 'V2E'", + ), + ( + """ + class C(NeighborConnectivity[Edge, Edge]): + Local = V2E.Local + """, + "cannot share the local dimension of 'V2E'", + ), + ( + """ + class C(NeighborConnectivity): + class Local(LocalDimensionIndex): ... + """, + "must derive from 'NeighborConnectivity\\[Domain, Codomain\\]'", + ), + ( + """ + class C(V2E): + class Local(LocalDimensionIndex): ... + """, + "must derive from 'NeighborConnectivity\\[Domain, Codomain\\]'", + ), + ( + """ + class C(NeighborConnectivity[V2E.Local, Edge]): + class Local(LocalDimensionIndex): ... + """, + "'Domain' 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 L(LocalDimensionIndex, kind=DimensionKind.HORIZONTAL): ... + """, + "is a local dimension and has no kind", + ), + ( + """ + class L(LocalDimensionIndex, size=1.5): ... + """, + "must be an integer", + ), + ], + ) + 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: typing.TypeAlias = Coeff + """ + ) + 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: typing.TypeAlias = 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_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] + + 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"): + + class C(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = 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, +) -> NeighborTableType: + """The type of a table no declaration names, as `NeighborTable.__gt_type__()` describes it.""" + structure = common.ConnectivityType( + domain=domain, codomain=codomain, skip_value=skip_value, dtype=core_defs.dtype(dtype) + ) + return NeighborTableType( + connectivity=structure, + skip_value=skip_value, + dtype=core_defs.dtype(dtype), + max_neighbors=max_neighbors, + ) + + +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_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( + """ + 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 TestNeighborTableType: + @staticmethod + def _v2e_shaped_table(codomain=Edge): + from gt4py.next import constructors + + return constructors.as_connectivity( + domain={Vertex: 2, V2E.Local: 4}, + codomain=codomain, + data=np.array([[0, 1, 2, -1], [1, 2, 3, 0]]), + skip_value=common._DEFAULT_SKIP_VALUE, + ) + + def test_domain_and_codomain_come_from_the_declaration(self): + table_type = common.check_neighbor_table(V2E, self._v2e_shaped_table()) + assert table_type.connectivity is V2E + assert table_type.domain == (Vertex, V2E.Local) + assert table_type.codomain is Edge + assert table_type.max_neighbors == 4 and table_type.has_skip_values + + def test_a_table_alone_has_its_structural_type(self): + structure = self._v2e_shaped_table().__gt_type__() + assert type(structure) is common.ConnectivityType + assert structure.domain == (Vertex, V2E.Local) and structure.codomain is Edge + + def test_bound_by_the_provider_key(self): + # a sharer's table looks like its owner's but for the codomain: only the key tells + # which declaration it is bound to + sharer = _declare( + """ + class V2V(NeighborConnectivity[Vertex, Vertex]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2V"] + v2e_table, v2v_table = self._v2e_shaped_table(), self._v2e_shaped_table(Vertex) + table_types = common.offset_provider_to_type( + {V2E.offset_tag: v2e_table, sharer.offset_tag: v2v_table, "undeclared": v2e_table} + ) + assert table_types[V2E.offset_tag].connectivity is V2E + assert table_types[sharer.offset_tag].connectivity is sharer + assert table_types[sharer.offset_tag].domain == table_types[V2E.offset_tag].domain + assert table_types[sharer.offset_tag].codomain is Vertex + # a key no declaration answers to leaves the table typed by its structure + assert table_types["undeclared"].connectivity == v2e_table.__gt_type__() + assert table_types["undeclared"].domain == table_types[V2E.offset_tag].domain + + def test_a_table_that_does_not_match_its_key(self): + with pytest.raises(ValueError, match="its codomain is"): + common.offset_provider_to_type({V2E.offset_tag: self._v2e_shaped_table(Vertex)}) + + def test_a_type_bound_to_another_declaration(self): + sharer = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2EShared"] + v2e_type = common.check_neighbor_table(V2E, self._v2e_shaped_table()) + with pytest.raises(ValueError, match="bound to 'V2E'"): + common.check_neighbor_table(sharer, v2e_type) + + def test_fingerprint_tells_the_declarations_apart(self): + from gt4py.next import fingerprinting + + sharer = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2EShared"] + table = self._v2e_shaped_table() + owner_type = common.check_neighbor_table(V2E, table) + sharer_type = common.check_neighbor_table(sharer, table) + # lenient: `_declare` classes are not importable + assert fingerprinting.lenient_fingerprinter( + owner_type + ) != fingerprinting.lenient_fingerprinter(sharer_type) + assert fingerprinting.lenient_fingerprinter( + owner_type + ) == fingerprinting.lenient_fingerprinter(common.check_neighbor_table(V2E, table)) + + +class TestFrontendIntegration: + def test_from_value_is_an_offset(self): + # NOTE: pins the `__gt_type__` branch of `from_value` ahead of the dimension branch; a + # connectivity declaration is a class, like a dimension. + assert type_translation.from_value(V2E) == ts.ShiftType( + codomain=Edge, domain=(Vertex, V2E.Local), tag=V2E.Local.tag + ) + + def test_shift_type_str(self): + assert str(V2E.__gt_type__()) == ( + f"Shift[{V2E.Local.tag}: {Edge} -> ({Vertex}, {V2E.Local})]" + ) + assert str(ts.ShiftType(codomain=KDim, domain=(KDim,))) == f"Shift[{KDim} -> {KDim}]" + + def test_field_offset_is_derived_once(self): + assert V2E.__gt_field_offset__() is V2E.__gt_field_offset__() + 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 domain_of(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return a(V2E.domain) + + # NOTE: `V2E` is typed as a `ts.ShiftType`, whose `domain` is a field of the type, not a + # value in DSL code + with pytest.raises(errors.DSLError, match="has no attribute 'domain'"): + FieldOperatorParser.apply_to_function(domain_of) + + def test_fingerprint_covers_the_declaration(self): + from gt4py.next import fingerprinting + + 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 + ) + 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) + + +class TestConnectivityKeyOver: + def _type(self, connectivity): + return _table_type(domain=(connectivity.domain, 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( + """ + class C(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = ConstList + """ + ) + + +class TestLocality: + """`DimensionKind.LOCAL` is gone: localness is the class, and a local dimension has no kind.""" + + def test_no_local_kind(self): + assert set(DimensionKind.__members__) == {"HORIZONTAL", "VERTICAL"} + + @pytest.mark.parametrize( + "dim, expected", + [ + (V2E.Local, True), + (LsqCoeff, True), + (common.ConstList, True), + (Vertex, False), + (KDim, False), + ], + ) + def test_is_local_dimension(self, dim, expected): + assert common.is_local_dimension(dim) is expected + + @pytest.mark.parametrize("value", [None, "V2E", V2E, LocalDimensionIndex(0)]) + def test_is_local_dimension_of_a_non_dimension(self, value): + assert common.is_local_dimension(value) is False + + @pytest.mark.parametrize("local", [V2E.Local, LsqCoeff, LocalDimensionIndex]) + def test_a_local_dimension_cannot_be_staggered(self, local): + # what keeps `is_local_dimension` total: there is no staggered local dimension + with pytest.raises(TypeError, match="not a declared Cartesian axis"): + common.Staggered[local] + + def test_display(self): + assert str(LsqCoeff) == "LsqCoeff[local]" + assert repr(LsqCoeff) == f"{LsqCoeff.tag}[local]" + + def test_layout_order_is_horizontal_local_vertical(self): + assert common.order_dimensions([KDim, LsqCoeff, Vertex]) == [Vertex, LsqCoeff, KDim] + with pytest.raises(ValueError, match="more than one local dimension"): + common.order_dimensions([Vertex, LsqCoeff, V2E.Local]) diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py index 2984713b76..f006282e53 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py @@ -14,6 +14,7 @@ Dimension, CartesianAxisIndex, DimensionIndex, + LocalDimensionIndex, DimensionKind, ) from gt4py.next.type_system import type_info, type_specifications as ts @@ -30,10 +31,10 @@ class JDim(CartesianAxisIndex): ... class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... class TDim(CartesianAxisIndex): ... diff --git a/typing_tests/pyright_probes.py b/typing_tests/pyright_probes.py new file mode 100644 index 0000000000..b89ffe6705 --- /dev/null +++ b/typing_tests/pyright_probes.py @@ -0,0 +1,144 @@ +# 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 0030). 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.CartesianAxisIndex, 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) + + +# -- Cartesian axis levels (ADR 0029). Each rejection carries a targeted ignore, and +# `reportUnnecessaryTypeIgnoreComment` turns any rejection that stops firing into an error. + + +def any_dimension(dim: gtx.Dimension) -> None: ... + + +def any_axis(dim: type[gtx.AnyCartesianAxisIndex]) -> None: ... + + +def declared_axis(dim: type[gtx.CartesianAxisIndex]) -> None: ... + + +def staggered_field(a: gtx.Field[gtx.Dims[Cell, gtx.Staggered[KDim]], gtx.float64]) -> None: ... + + +any_dimension(gtx.Staggered[KDim]) # a staggered dimension is still a dimension +any_dimension(V2E.Local) # local dimensions keep their place below the root +any_axis(gtx.Staggered[KDim]) +declared_axis(KDim) +_shift_staggered = gtx.Staggered[KDim] + 1 +_shift_half = KDim + 0.5 + +any_axis(Cell) # pyright: ignore[reportArgumentType] +any_axis(V2E.Local) # pyright: ignore[reportArgumentType] +declared_axis(gtx.Staggered[KDim]) # pyright: ignore[reportArgumentType] +_doubly: typing.TypeAlias = gtx.Staggered[gtx.Staggered[KDim]] # pyright: ignore[reportInvalidTypeArguments] +_location: typing.TypeAlias = gtx.Staggered[Cell] # pyright: ignore[reportInvalidTypeArguments] +_local: typing.TypeAlias = gtx.Staggered[V2E.Local] # pyright: ignore[reportInvalidTypeArguments] +_shift_location = Cell + 1 # pyright: ignore[reportOperatorIssue] +_shift_local = V2E.Local - 1 # pyright: ignore[reportOperatorIssue] diff --git a/typing_tests/pyrightconfig.json b/typing_tests/pyrightconfig.json new file mode 100644 index 0000000000..7dc8caf8d9 --- /dev/null +++ b/typing_tests/pyrightconfig.json @@ -0,0 +1,6 @@ +{ + "typeCheckingMode": "standard", + "reportMissingImports": "error", + "reportMissingTypeStubs": "none", + "reportUnnecessaryTypeIgnoreComment": "error" +} diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 1506c46a75..a65d3e7d09 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -302,3 +302,74 @@ main:16:48: error: Type argument "C" of "Staggered" must be a subtype of "CartesianAxisIndex" [type-var] main:17:13: error: Unsupported operand types for + ("type[C]" and "int") [operator] main:18:18: error: Unsupported operand types for - ("type[C]" and "int") [operator] + + - case: neighbor_connectivity_declaration + main: | + from __future__ import annotations + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6): + class Local(gtx.LocalDimensionIndex): ... + + def sparse(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + reveal_type(V2E.Local) + reveal_type(V2E.Local.owner) + reveal_type(V2E[1]) + out: | + main:12:13: note: Revealed type is "def (value: int) -> main.V2E.Local" + main:13:13: note: Revealed type is "type[gt4py.next.common.NeighborConnectivity[Any, Any]] | None" + main:14:13: note: Revealed type is "gt4py.next.common.Connectivity[Any, Any]" + + - case: neighbor_connectivity_locals_are_distinct + main: | + from __future__ import annotations + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + class Cell(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + class V2C(gtx.NeighborConnectivity[Vertex, Cell]): + class Local(gtx.LocalDimensionIndex): ... + + def takes_v2e(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + def caller(a: gtx.Field[gtx.Dims[Vertex, V2C.Local], gtx.float64]) -> None: + takes_v2e(a) + out: | + main:17:15: error: Argument 1 to "takes_v2e" has incompatible type "Field[Dims[Vertex, main.V2C.Local], float]"; expected "Field[Dims[Vertex, main.V2E.Local], float]" [arg-type] + main:17:15: note: "Field[Dims[Vertex, Local], float].__call__" has type "def __call__(self, index_field: Connectivity[Any, Any] | FieldOffset | type[NeighborConnectivity[Any, Any]], *args: Connectivity[Any, Any] | FieldOffset | type[NeighborConnectivity[Any, Any]]) -> Field[Any, Any]" + + - case: neighbor_connectivity_generic_local + main: | + from __future__ import annotations + import typing + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + L = typing.TypeVar("L", bound=gtx.LocalDimensionIndex) + + def local_of(conn: type[gtx.NeighborConnectivity]) -> type[gtx.LocalDimensionIndex]: + # `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 + + def caller(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: + reveal_type(first(a)) + out: | + main:22:17: note: Revealed type is "type[main.V2E.Local]" diff --git a/uv.lock b/uv.lock index ef2d5ca914..648636a40e 100644 --- a/uv.lock +++ b/uv.lock @@ -1442,6 +1442,7 @@ typing = [ ] typing-exports = [ { name = "mypy", extra = ["faster-cache"] }, + { name = "pyright" }, { name = "pytest-mypy-plugins" }, { name = "types-decorator" }, { name = "types-docutils" }, @@ -1598,6 +1599,7 @@ typing = [ ] typing-exports = [ { name = "mypy", extras = ["faster-cache"], specifier = ">=1.13.0" }, + { name = "pyright", specifier = ">=1.1.400" }, { name = "pytest-mypy-plugins", specifier = ">=4.0.0" }, { name = "types-decorator", specifier = ">=5.1.8" }, { name = "types-docutils", specifier = ">=0.21.0" }, @@ -3126,6 +3128,19 @@ version = "2.1" source = { registry = "https://pypi.org/simple" } sdist = { url = "https://files.pythonhosted.org/packages/bc/7c/d724ef1ec3ab2125f38a1d53285745445ec4a8f19b9bb0761b4064316679/pyreadline-2.1.zip", hash = "sha256:4530592fc2e85b25b1a9f79664433da09237c1a270e4d78ea5aa3a2c7229e2d1", size = 109189, upload-time = "2015-09-16T08:24:48.745Z" } +[[package]] +name = "pyright" +version = "1.1.414" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nodeenv" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/e1/1b/244c7b710031ada80f27e579ec20d28a2285dfc318fed0339866b1047f12/pyright-1.1.414.tar.gz", hash = "sha256:523c0a97c60da6333234955c277730c9cf4f5bd6d5399e7b7d2b0fc5d3599524", size = 4154638, upload-time = "2026-09-10T12:26:53.181Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d7/ba/18b6e682ead424ad24bcc134339ae5d1b931cd9ae260540592a058a91279/pyright-1.1.414-py3-none-any.whl", hash = "sha256:2a6b4b3298c9eec174c5ed83bd338de6eee82df2992f3e1930e6199d381be36f", size = 6225049, upload-time = "2026-09-10T12:26:51.427Z" }, +] + [[package]] name = "pytest" version = "9.1.1"