From 565fa7558f6ed5a2e14c4b0b646803b62a784443 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 18 Sep 2026 04:32:21 +0200 Subject: [PATCH 01/20] docs[next]: ADR 0028, dimensions as nominal types Records the design this stack implements, ahead of the code, because it reverses decisions that are easier to argue about as prose than as a 150-file diff. A concrete dimension becomes a class and an index an instance of it, so `Field[Dims[IDim], float64]` type-checks with no gt4py mypy plugin. A dimension's identity is the Python type and its tag is the qualified Python name, which is also a unique, valid IR spelling. The decisions worth reviewing first, each with the reason it is not the obvious choice: * Nominal identity, not `(name, kind)` value equality with an interning registry. The registry decouples the Python type's identity from the IR's and needs `copyreg` plus a custom fingerprint deconstructor to bridge the gap. Under value equality the `typing` subscription cache already aliases `Field[Dims[I]]` and `Field[Dims[I2]]` for two distinct same-named classes, so the static and runtime views disagree exactly there. * `resolve(tag)` is an import, so types reaching the IR must be declared at module level. Interactive `__main__` is a documented limitation; `spawn` workers re-execute the main script, so file-based `__main__` resolves. * Generated identifiers need a *prefix* escape (`_` -> `_u`, `.` -> `_d`). The obvious `_` -> `__` then `.` -> `_` is not injective: a dot becomes a single underscore, so `".."` collides with an escaped `"_"`. * `Staggered[D]` supersedes ADR 0026's `_Staggered` name prefix, which cannot survive type identity. It cannot be a PEP 695 generic either -- that yields a `_GenericAlias`, not a class -- so it is an interning metaclass paired with a `TYPE_CHECKING` declaration. * `DimensionMeta` must declare `__hash__` explicitly, since `__eq__` stays for the `I == 5` overload and Python would otherwise make every dimension class unhashable. Consequence recorded for ADR 0023: a dimension is fingerprinted by qualified name, so moving a declaration between modules invalidates compiled artifacts. --- .../next/0028-Dimensions_As_Nominal_Types.md | 233 ++++++++++++++++++ docs/development/ADRs/next/README.md | 1 + 2 files changed, 234 insertions(+) create mode 100644 docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md new file mode 100644 index 0000000000..a62f400d3b --- /dev/null +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -0,0 +1,233 @@ +--- +tags: [] +--- + +# Dimensions as Nominal Types + +- **Status**: proposed +- **Authors**: Enrique González Paredes (@egparedes) +- **Created**: 2026-09-18 +- **Updated**: 2026-09-18 + +A concrete dimension becomes a **class**, and an index along it an **instance** of +that class — the shape `enum.Enum` uses, where the class is the collection and +the instances are its members: + +```python +class IDim(gtx.DimensionIndex): ... + + +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +IDim # the dimension -- annotated `gtx.Dimension` +IDim(0) # an index into it -- annotated `IDim` +``` + +so `gtx.Field[gtx.Dims[IDim], gtx.float64]` is valid for any PEP 484 checker with +no gt4py mypy plugin. + +A dimension's **identity is the Python type**, and its `tag` — the string that +crosses into the IR and the generated code — is its **qualified Python name**, +`f"{cls.__module__}.{cls.__qualname__}"`. + +## Context + +`Dimension` was a frozen dataclass whose instances are values, not types. Two +consequences drove this change. + +**Dimensions were not usable as types.** `Field[[IDim], float64]` needed a +dedicated mypy plugin that substituted at most four distinct placeholders per run +(`_DimA`..`_DimD`, then `_AnyDim` for everything after), made `TypeVar`s over +dimensions impossible, and served only mypy — no other checker. See #2503. + +**Names are load-bearing in too many places.** A neighbor connectivity is +currently spread over four independently authored strings that must agree by +string equality and are never checked against each other at declaration time: the +`FieldOffset` tag, the Python variable it is bound to, the local `Dimension`'s +name, and the `offset_provider` key. Whichever one reaches +`common.get_offset` depends on the execution path and the operation. Making a +dimension's identity its Python type is the prerequisite for collapsing those +names into one declaration (a follow-up ADR covers the connectivity half). + +## Decision + +### Identity is the Python type; the tag is the qualified name + +Two dimension classes are the same dimension if and only if they are the same +class. The static view (checkers see nominal types) and the runtime view +(equality is `is`) agree by construction, and the tag is a unique string that is +also a valid IR spelling. + +The alternative — `(name, kind)` value equality plus an interning registry, so +that independently declared same-named dimensions stay interchangeable — was +considered and rejected. It decouples the Python type's identity from the IR's, +and needs a registry, a `copyreg` hook and a custom fingerprint deconstructor to +paper over that gap. It also cannot give a dimension nested inside another +declaration a unique name without further convention, which the connectivity +work requires. + +One concrete argument in favour of nominal identity: under `(name, kind)` +equality the `typing` subscription cache aliases `Field[Dims[I]]` and +`Field[Dims[I2]]` for two *distinct* same-named classes, so the static and +runtime views disagree exactly there. Under type identity that aliasing +disappears. + +### Consequences of the tag being a qualified name + +1. **Reconstruction from the IR is an import.** `common.resolve(tag)` imports the + module and walks the qualname; nested declarations resolve naturally. The IR + references a Python type exactly the way `pickle` references a class. It is + memoized, because type inference calls it once per `AxisLiteral`. + + A purely dotted tag does not record where the module path ends and the + qualname begins, so `resolve` tries the *longest importable prefix* and walks + the remainder. A collision requires a module path and an attribute chain to + have the same spelling; a real module always wins. `pickle` avoids the + ambiguity by storing the two parts separately, and that remains available if + the residual ever bites. + +2. **Types reaching the IR must be importable**, i.e. declared at module level. + `__init_subclass__` rejects a `` qualname as an early heuristic; it is + neither necessary nor sufficient (`type("Dyn", ...)` inside a function passes, + a class deleted after creation passes), so the authoritative check remains + pickle's own `save_global`. + + **Known limitation**: interactive `__main__` — the REPL, notebooks, + `python -c` — cannot be resolved. The `spawn`-based compile workers + re-execute the main *script* as `__mp_main__`, so a dimension declared in a + file's `__main__` does resolve, provided the script has the + `if __name__ == "__main__":` guard the worker pool already requires. + +3. **No registry and no blanket `copyreg`.** Module-level classes pickle by + reference, which is correct. A *narrow* `copyreg` registration is still + required for parametrized dimensions — see Staggered below. + +4. **Cache fingerprints depend on module paths.** A dimension is fingerprinted by + qualified name, so moving a declaration between modules invalidates compiled + artifacts. This is a consequence for the build cache of ADR 0023, not a + reversal of it. The generic `type` deconstructor is correct for the *lenient* + fingerprint variant; the STRICT variant rejects a parametrized dimension, + which is not importable under its qualified name. + +5. **Generated identifiers need injective mangling.** A dot is illegal in a C++ + identifier, in a DaCe symbol, and in `eve`'s `SymbolName` + (`^[a-zA-Z_]\w*$`). One shared pair, used by every backend: + + ```python + def codegen_name(tag: Tag) -> str: + return tag.replace("_", "_u").replace(".", "_d") + + + def from_codegen_name(name: str) -> Tag: + return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) + ``` + + A *prefix* escape, not `_ -> __` followed by `. -> _`: the latter is **not + injective**, since a dot becomes a single underscore and `".."` collides with + an escaped `"_"`. Every `_` in the output is the first character of a + two-character escape, so decoding is unambiguous. Names grow, which is what + gtfn's existing `TagDefinition.alias` mechanism is for. + +6. **Staggered dimensions become a real parametrized type.** ADR 0026's + `_Staggered` *name prefix* cannot survive type identity: `Dimension(f"_Staggered{name}")` + names no importable type, and the prefix cannot recover the base dimension's + module. `Staggered[D]` replaces it and supersedes that part of ADR 0026. + + A PEP 695 generic does not work: `Staggered[KDim]` would be a + `typing._GenericAlias`, not a class, so it fails `issubclass` and eve's + `type[DimensionIndex]` validation, and its tag cannot name the base. Instead a + metaclass `__getitem__` builds and **interns a real class**, paired with a + `TYPE_CHECKING` declaration so checkers still see an ordinary generic: + + - bases are `(Staggered,)` and deliberately **not** `(Staggered, base)`: + a staggered dimension is a *different* dimension, so + `issubclass(Staggered[KDim], KDim)` must be false. Only `kind` is inherited. + - `is_staggered(dim)` is `"base" in dim.__dict__` and + `as_non_staggered(dim)` is `dim.base` — structural, no string sniffing. + `issubclass(dim, Staggered)` would be wrong, because it is also true of the + bare base, which has no `base`. + - `resolve` gains a `[]` grammar and evaluates + `Staggered[resolve(inner)]`, hitting the same intern table, so a staggered + dimension round-trips through the IR to the *same* class object. This is the + one place where "resolution is an import" is not literally true. + - the intern table is keyed by a dimension *class*, not by a user-authored + name. It is memoization of a type constructor, as `typing`'s own + subscription cache is — not the name-keyed registry this ADR rejects. + - `Staggered[KDim].__qualname__` contains brackets, which + `pickle.save_global` cannot look up, so `copyreg` is registered on the + staggered metaclass. It must fall back to by-reference pickling for the bare + base, which is also an instance of that metaclass. + +### `Dimension` is annotation-only + +`common.Dimension` becomes a PEP 695 alias for `type[DimensionIndex]`, so the +removed `gtx.Dimension("I")` spelling raises rather than misbehaving: a plain +`TypeAlias` for `type[X]` is a `types.GenericAlias`, and calling one forwards to +`__origin__` while discarding the arguments — it would evaluate to `str` with no +error. A `TypeAliasType` is simply not callable. + +Its cost is that `get_origin()` of such an alias is `None`, so a site dispatching +on an annotation's shape must resolve it first (`xtyping.resolve_annotation`, +added in #2841). + +### Naming + +A dimension's name is `.tag` (typed `common.Tag`, which already existed) and +`.value` keeps its meaning as the index position, on the *instance*. The reverse +split does not type-check at all — an instance attribute cannot shadow a +`ClassVar` — and this direction leaves every index expression untouched. +Reading `.value` on a dimension *class* raises a metaclass `AttributeError` +pointing at `.tag`, rather than returning the `__slots__` member descriptor and +surfacing much later as a missing offset-provider key. + +`tag` is a metaclass **property**, so it cannot drift from the type. A class-body +`tag = "..."` would therefore be silently ignored, which is exactly the renaming +pattern downstream code uses — so `__init_subclass__` raises on it. + +`DimensionMeta` must declare `__hash__ = type.__hash__` explicitly: Python sets +`__hash__ = None` on any class body defining `__eq__` without it, and `__eq__` +stays for the `I == 5` → `Domain` overload. Without it every dimension class is +unhashable and `ts.DimensionType` fails at *import*. + +## Consequences + +- `common.NamedIndex` is deleted; `.dim` and `.value` keep working, on the index + instance. +- The dimension half of `mypy_plugin.py` is deleted; only the mixed-precision + hooks remain. Static checking is now available to pyright as well as mypy. +- Every declaration in the tree — including docs, workshop notebooks and + `examples/`, which `test_examples` executes — becomes a class statement. There + is no `dimension(tag, kind)` factory for user code, so there is no minimal + migration form. +- Test modules that declared dimensions inside test functions must move them to + module level. Where two same-named function-local dimensions were silently the + same dimension, they are now distinct — each resulting failure is a real + finding. +- `repr()` of a dimension is `I[horizontal]`; `str()` is unchanged, so error + messages are byte-identical. + +## Alternatives considered + +- **`(name, kind)` value equality with an interning registry.** See Decision. The + test-tree consequence of rejecting it is real: many function-local dummy + dimensions must move to module level. +- **Keeping the mypy plugin.** Serves one checker, caps distinct dimensions at + four per run, and blocks `TypeVar`s over dimensions. +- **`AxisLiteral` carrying the dimension class** instead of `value: str`. Would + remove `resolve` from the type-inference hot path, but changes IR node shape in + the same step as the dimension rewrite. Deferred. +- **`Staggered` as a PEP 695 generic.** Does not produce a class; see Decision 6. + +## References + +- Implements the `shared/dimensions-as-types` proposal (gt4py_knowledge#27, + @havogt) and the `egparedes/connectivities-as-types` proposal + (gt4py_knowledge#32). +- Closes the static-typing gap reported in #2503. +- Supersedes the `_Staggered` name-prefix mechanism of + [ADR 0026](0026-Staggered_Dimensions.md); the indexing convention there is + unchanged. +- Consequence for the build cache of [ADR 0023](0023-Fingerprinting.md). +- An alternative to #2844, which implements the same class-shaped dimension with + value identity. diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index f19167ef9b..fd74e97eb4 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -22,6 +22,7 @@ Writing a new ADR is simple: - [0021 - Argument Descriptors](0021-Argument-Descriptors.md) - [0023 - Fingerprinting](0023-Fingerprinting.md) - [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md) +- [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md) ### Frontend and Parsing #frontend From 823c0dda7bffdbede166584d681032f5b6525ed2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 18 Sep 2026 07:53:41 +0200 Subject: [PATCH 02/20] feat[next]: add injective tag mangling for generated identifiers A dimension tag becomes a qualified Python name in this stack, so it contains dots -- illegal in a C++ identifier, in a DaCe symbol, and in `eve`'s `SymbolName` (`^[a-zA-Z_]\w*$`). `codegen_name` mangles a tag into a valid identifier and `from_codegen_name` recovers it, for the backends that parse a generated name back into the dimension it refers to. The escape is a *prefix* escape (`_` -> `_u`, `.` -> `_d`). The obvious alternative -- double every underscore, then turn dots into single underscores -- is **not injective**: a dot becomes a single underscore, so `'..'` and `'_'` both map to `'__'`. Exhaustively: 686 collisions in 1092 inputs over `{a, ., _}` up to length 6. A collision here would mean two distinct dimensions silently sharing one generated symbol, i.e. wrong results rather than a crash, so the test is exhaustive rather than example-based: every string over `{a, ., _, u, d}` up to length 5 mangles uniquely and round-trips, and the adversarial cases that look like the escape sequences themselves (`_u`, `_d`, `a_ud.b`) are pinned separately. No caller yet; the sites that need it arrive with the dimension classes. --- src/gt4py/next/common.py | 51 ++++++++++++++++++++++ tests/next_tests/unit_tests/test_common.py | 49 +++++++++++++++++++++ 2 files changed, 100 insertions(+) diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 5bc53f6474..eaf8aa8464 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -14,6 +14,7 @@ import enum import functools import math +import re import sys import types from collections.abc import Callable, Iterable, Mapping, Sequence @@ -60,6 +61,56 @@ class Dims(tuple[Unpack[ShapeTs]]): ... Tag: TypeAlias = str +def codegen_name(tag: Tag) -> str: + """ + Mangle a dimension or offset tag into a valid generated identifier. + + A tag is a qualified Python name, so it contains dots, which are illegal in a C++ + identifier, in a DaCe symbol, and in `eve`'s `SymbolName` (`^[a-zA-Z_]\\w*$`). Since a + generated identifier may only contain `[A-Za-z0-9_]`, the underscore is the only + available separator, and escaping it is what makes the mapping reversible. + + The escape is a *prefix* escape. The obvious alternative -- double every underscore, + then turn dots into single underscores -- is **not injective**: a dot becomes a single + underscore, so `'..'` and `'_'` both map to `'__'`. + + Args: + tag: A dimension or offset tag, i.e. a qualified Python name. + + Returns: + A valid identifier, unique for each distinct `tag`. + + Examples: + >>> codegen_name("mod.V2E.Local") + 'mod_dV2E_dLocal' + >>> codegen_name("a_b.c") + 'a_ub_dc' + >>> from_codegen_name(codegen_name("my__mod.X")) + 'my__mod.X' + """ + return tag.replace("_", "_u").replace(".", "_d") + + +def from_codegen_name(name: str) -> Tag: + """ + Recover a tag from the identifier `codegen_name` produced for it. + + Needed wherever a backend parses a generated name back into the dimension or offset it + refers to. + + Args: + name: An identifier produced by `codegen_name`. + + Returns: + The original tag. + + Examples: + >>> from_codegen_name("mod_dV2E_dLocal") + 'mod.V2E.Local' + """ + return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) + + @enum.unique class DimensionKind(StrEnum): HORIZONTAL = "horizontal" diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index fddc9dbd2e..84fae1145e 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -6,6 +6,7 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import itertools import operator from typing import Optional, Pattern @@ -793,3 +794,51 @@ def test_different_dims_raises(self): d2 = Domain(dims=(JDim,), ranges=(UnitRange(3, 10),)) with pytest.raises(NotImplementedError, match="different dimensions"): d1 | d2 + + +class TestCodegenName: + """`codegen_name` must be injective: a collision means two dimensions sharing a symbol.""" + + @pytest.mark.parametrize( + "tag, expected", + [ + ("IDim", "IDim"), + ("mod.V2E.Local", "mod_dV2E_dLocal"), + ("a.b_c", "a_db_uc"), + ("a_b.c", "a_ub_dc"), + ("my__mod.X", "my_u_umod_dX"), + ("_CONST_DIM", "_uCONST_uDIM"), + ], + ) + def test_known_values(self, tag, expected): + assert common.codegen_name(tag) == expected + assert common.from_codegen_name(expected) == tag + + @pytest.mark.parametrize("tag", ["_u", "_d", "a_ud.b", "_ud_du", "..", "__"]) + def test_roundtrip_adversarial(self, tag): + """Tags that look like the escape sequences themselves must still round-trip.""" + assert common.from_codegen_name(common.codegen_name(tag)) == tag + + def test_output_is_a_valid_identifier(self): + for tag in ["mod.V2E.Local", "a.b_c", "_CONST_DIM", "pkg.sub.Dim"]: + assert re.fullmatch(r"[A-Za-z_]\w*", common.codegen_name(tag)), tag + + def test_injective_and_reversible_exhaustively(self): + """ + Exhaustive over the characters that can actually collide, to a length that covers + every interaction between them. + + The naive scheme -- `_` -> `__` then `.` -> `_` -- fails this with 686 collisions, + because a dot becomes a single underscore and `'..'` collides with an escaped `'_'`. + """ + alphabet = "a._ud" + seen: dict[str, str] = {} + for length in range(1, 6): + for tag in map("".join, itertools.product(alphabet, repeat=length)): + mangled = common.codegen_name(tag) + assert mangled not in seen, ( + f"collision: {seen.get(mangled)!r} and {tag!r} both map to {mangled!r}" + ) + seen[mangled] = tag + assert common.from_codegen_name(mangled) == tag + assert len(seen) == sum(len(alphabet) ** n for n in range(1, 6)) From 7ea399a778a4eb9086dd3776679895f076b7f038 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 18 Sep 2026 08:00:16 +0200 Subject: [PATCH 03/20] wip[next]: dimension classes, resolve() and Staggered[D] in common The core of ADR 0028. NOT green: `Dimension("X")` now raises, so every declaration in the tree has to become a class statement. That sweep is the next commit; this one is the mechanism it depends on. * `DimensionMeta` / `DimensionIndex`: a dimension is a class, an index an instance. Identity is the type; `tag` is the qualified Python name and is a metaclass *property*, so it cannot drift from what it names. `__hash__` is declared explicitly because `__eq__` stays for the `I == 5` overload. Dimension-vs-dimension `__eq__` is deliberately *not* overridden -- identity is the correct answer, and overriding it is what the ADR rejects. * `type Dimension = type[DimensionIndex]`, a PEP 695 alias so the removed `Dimension("I")` raises instead of silently evaluating to `str`. * Display uses `__qualname__`, not `tag`: 148 tests assert on message text like `Field[[IDim], float64]`, and a qualified name there is noise. `repr` carries the module. * `resolve(tag)` imports and walks the qualname, memoized, with a grammar for the parametrized `[]` form. * `Staggered[D]` replaces ADR 0026's `_Staggered` name prefix, which cannot survive type identity. An interning metaclass builds a *real* class, paired with a `TYPE_CHECKING` declaration -- a PEP 695 generic yields a `_GenericAlias`, which is not a class and fails eve's `type[...]` validation. Bases are `(Staggered,)` and deliberately not `(Staggered, base)`, so a staggered field is not accepted where its base is required. Verified: real class, interning stable, kind inherited, tag resolvable, instances work, `copyreg` round-trips both the parametrized class and the bare base, and all three escape routes (subclassing either level, double subscript) are blocked. * `ConstListDim` is declared **once** in `common`. It used to be built independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless while dimensions compared by `(name, kind)`. Under nominal identity those would be two different dimensions and the `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s built by embedded execution. --- .../next/0028-Dimensions_As_Nominal_Types.md | 10 + ...ectivities-as-types-implementation-plan.md | 1033 +++++++++++++++++ .../next/fieldoffset-tag-constraints.md | 551 +++++++++ src/gt4py/next/common.py | 413 ++++++- src/gt4py/next/iterator/embedded.py | 2 +- .../dace/lowering/gtir_to_sdfg_lambda.py | 5 +- 6 files changed, 1956 insertions(+), 58 deletions(-) create mode 100644 docs/development/next/connectivities-as-types-implementation-plan.md create mode 100644 docs/development/next/fieldoffset-tag-constraints.md diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index a62f400d3b..08af4b4d24 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -185,6 +185,16 @@ surfacing much later as a missing offset-provider key. `tag = "..."` would therefore be silently ignored, which is exactly the renaming pattern downstream code uses — so `__init_subclass__` raises on it. +### Display uses the unqualified name + +The `tag` is qualified, but it is *identity and IR spelling*, not a display name. +User-facing diagnostics use `cls.__qualname__`, so `Field[[IDim], float64]` reads +the same as before instead of becoming `Field[[pkg.mod.IDim], float64]`. This is +not a third name concept: `__qualname__` is a Python builtin attribute, and for a +nested declaration it is already the readable form (`V2E.Local`). `repr()` shows +the full tag and the kind, which disambiguates the rare case of two same-named +dimensions from different modules appearing in one message. + `DimensionMeta` must declare `__hash__ = type.__hash__` explicitly: Python sets `__hash__ = None` on any class body defining `__eq__` without it, and `__eq__` stays for the `I == 5` → `Domain` overload. Without it every dimension class is diff --git a/docs/development/next/connectivities-as-types-implementation-plan.md b/docs/development/next/connectivities-as-types-implementation-plan.md new file mode 100644 index 0000000000..52bdbbc0c1 --- /dev/null +++ b/docs/development/next/connectivities-as-types-implementation-plan.md @@ -0,0 +1,1033 @@ +# Connectivities as types — implementation plan + +**Status**: **APPROVED** at revision 6 (adversarial review rounds 1–4; round 4 verdict APPROVED) +**Target**: an 8-PR stack on `main`, *alternative to* GridTools/gt4py#2844 +**Proposal**: `egparedes/connectivities-as-types` in GridTools/gt4py_knowledge (PR #32) +**Baseline tree**: `b3c53fa7e` (v1.2.2) + +## 0. Scope and relation to #2844 + +The proposal and #2844 agree on *what a dimension is* (a class, its indices its +instances) and disagree on *what identity a dimension has*. #2844 chose +`(tag, kind)` value equality with an interning registry; the proposal chose +nominal type identity with the tag being the qualified Python name. That single +disagreement propagates into five mechanisms, so the two cannot both land. + +This stack **re-cuts** #2844: it keeps the machinery independent of identity, +drops the machinery that exists only to support value identity, and then builds +the connectivity layer on top. **#2844 is closed, not merged** — that is what +makes this an alternative. Consequences: the new dimension ADR is **0028** (the +ADR directory on `main` ends at **0027**; 0028 exists only inside unmerged +#2844), and there is nothing to supersede. + +### Taken from #2844, unchanged in substance + +| Piece | Where in #2844 | +| --- | --- | +| `DimensionMeta` metaclass; `I + 1`, `I > 5`, `repr` living on it | `common.py` | +| `DimensionIndex` base: `__slots__ = ("value",)`, `kind` class keyword, `.dim` property | `common.py` | +| `type Dimension = type[DimensionIndex]` as a PEP 695 alias (so `Dimension("I")` raises rather than silently evaluating to `str`) | `common.py` | +| Metaclass `.value` property raising `AttributeError` that points at `.tag` | `common.py` | +| Deletion of `common.NamedIndex` (`.dim` / `.value` move onto the index instance) | `common.py` + ~40 call sites | +| Deletion of the dimension half of `mypy_plugin.py` (`_DimA`..`_AnyDim`); only the mixed-precision hooks remain | `type_system/mypy_plugin.py` | +| The mechanical migration of every declaration, incl. docs, workshop notebooks and `examples/` (which `test_examples` executes) | 131 files: 52 `src/`, 66 `tests/`, 13 docs | +| `xtyping.resolve_annotation` usage at `fbuiltins._type_conversion_helper` (already on `main` via #2841) | — | + +### Dropped from #2844 + +| Piece | Why | +| --- | --- | +| `_DIMENSION_REGISTRY` interning | identity is the type; nothing to intern | +| `copyreg.pickle(DimensionMeta, _reduce_dimension)` — the **blanket** registration on all dimensions | verified: a module-level dimension class pickles by reference with no help (`pickle.loads(pickle.dumps(KDim)) is KDim`). **But a narrow `copyreg` on `StaggeredMeta` is still required** — see §1.5 | +| The `DimensionMeta`-vs-`DimensionMeta` branch of `__eq__` / `__ne__` | becomes `is`. **The `IntegralScalar` overloads (`I == 5` → `Domain`) are kept**, and therefore so is an explicit `__hash__` — see §1.0 | +| `DimensionIndex.__eq__` comparing `type(self) == type(other)` | becomes `type(self) is type(other)` | +| `common.dimension(tag, kind)` factory | replaced by `common.resolve(tag)`, which imports | +| `fingerprinting.py` `DimensionMeta` deconstructor keyed on `(tag, kind)` | under type identity a dimension *is* fingerprinted by qualified name, so the generic `type` deconstructor is correct — **for the lenient variant only**. The STRICT variant rejects `Staggered[KDim]`, which is not importable under its qualified name. Both in-tree fingerprinters are lenient (`ffront/stages.py:62`, `iterator/ir.py:26`), and `eve_utils.content_hash` (`compiled_program.py:420`) is pickle-based and so needs §1.5's `copyreg`. Record the STRICT caveat in the ADR | +| ADR 0028 as drafted in #2844 | never lands; this stack writes its own 0028 | + +### Changed relative to #2844 + +| Piece | #2844 | This stack | +| --- | --- | --- | +| `tag` default | `cls.__name__`, settable in the class body | `f"{cls.__module__}.{cls.__qualname__}"`, a metaclass property; a class-body `tag = ...` is a `TypeError` | +| Rebuilding a dimension from a tag | `dimension(tag, kind)` (registry) | `resolve(tag)` (`import_module` + `qualname` walk), memoized | +| Declaration site requirement | none | module level, or unpicklable; `` heuristic in `__init_subclass__` | +| Backend name mangling | `tag` used directly | `codegen_name(tag)` + inverse, at ~19 enumerated sites in two name spaces (§1.3(b), (c)) | +| Staggered dimensions | `_Staggered` prefix through the interning factory | `Staggered[D]`, a real parametrized type — **required in PR 2**, not optional (§1.5) | + +### Superseded + +- **#2845** (`test[next]: adopt class-style dimension declarations`) is subsumed + by PR 2: because `dimension()` is not user-facing, the minimal + `I = gtx.dimension("I")` form does not exist and every declaration takes class + form immediately. #2845's pyright coverage is folded in. +- The `FieldOffset`-as-frontend-identifier part of **ADR 0019**. +- **ADR 0026**'s `_Staggered` name prefix (PR 2). + +## 1. Design questions closed before implementation + +Everything in this section was verified by running it, not by reading. Probe +files are named; they become committed test material in the PR that needs them. + +### 1.0 Metaclass mechanics that are easy to get wrong + +**`__hash__` must be declared explicitly.** Python sets `__hash__ = None` on any +class body that defines `__eq__` without `__hash__` — metaclasses included. Since +the `I == 5` → `Domain` overload keeps `__eq__` on `DimensionMeta`, dropping +#2844's `__hash__` makes every dimension class *unhashable*: + +``` +>>> class M(type): +... def __eq__(cls, o): return True +>>> M.__hash__ is None +True +>>> class C(metaclass=M): pass +>>> hash(C) +TypeError: unhashable type: 'M' +``` + +That would break `domain({I: 2})` (`common.py:672-690`), +`Counter[common.Dimension]` (`embedded/nd_array_field.py:314`), +`dict[Dimension, SymbolicRange]` (`iterator/ir_utils/domain_utils.py:136,152`), +`seen: dict[Dimension, Dimension]` (`common.py:1351`), and eve's validator +memoization on annotation objects (`eve/type_validation.py:599`) — so +`ts.DimensionType` would fail at *import*. **Fix: `__hash__ = type.__hash__` +explicitly on `DimensionMeta`,** and likewise on `ConnectivityMeta` if it ever +defines `__eq__`. + +**A metaclass `__getitem__` shadows `__class_getitem__`.** `ConnectivityMeta` +needs `__getitem__` for `V2E[1]` (the single-neighbor shift handle that +`FieldOffset.__getitem__` provides today), but metaclass lookup takes precedence +over `Generic.__class_getitem__`, so a naive implementation makes +`NeighborConnectivity[V, E]` in a bases list fail with +`TypeError: tuple expected at most 1 argument, got 3`. + +**Fix, verified clean under `mypy --strict` and pyright 1.1.414 on Python 3.12** +(`/tmp/probe_meta_getitem3.py`): dispatch on the argument type, delegating +non-`int` subscription back to `cls.__class_getitem__`: + +```python +class ConnectivityMeta(type): + __hash__ = type.__hash__ + @overload + def __getitem__(cls, item: int) -> Connectivity: ... + @overload + def __getitem__(cls, item: Any) -> Any: ... + def __getitem__(cls, item: Any) -> Any: + # `numbers.Integral`, not `int`: `V2E[np.int32(1)]` must not fall through + # to the type-parameter branch (it raises `TypeError: V2E is not a + # generic class` there). `bool` is excluded so `V2E[True]` is an error + # rather than silently neighbor 1. + if isinstance(item, numbers.Integral) and not isinstance(item, bool): + return _bound_single_neighbor(cls, int(item)) + # type-parameter subscription, e.g. `NeighborConnectivity[V, E]` + return cast(Any, cls).__class_getitem__(item) +``` + +`cast(Any, cls)`, not `super()` — `__class_getitem__` is on the class, not on the +metaclass MRO; `super().__class_getitem__` raises `AttributeError`. With the cast +both checkers report zero errors and all uses work at runtime +(`NC[V, E]`, `class V2E(NC[V, E])`, `V2E[1]`, and `V2E.Local` as an annotation). +pyright accepts `NC[V, E]` in a **bases list**; in a *value* position it types it +`Any`, which is why the overloads above matter — without them `V2E[1]` is also +`Any` and the shift handle is untyped. + +### 1.1 `NeighborConnectivity` is **not** a `Connectivity` (proposal Open Q6) + +`common.Connectivity` is `Field[DimsT, IntegralScalar]` — a **data** protocol +(`common.py:990`; `ndarray`, `asnumpy`, `domain` are all on it). A declaration +class holds no data. + +**Resolution.** Two distinct things, distinct hierarchies: + +- `NeighborConnectivity` — a **declaration**. Not a `Connectivity`. It produces a + `NeighborConnectivityType` via `__gt_type__()`, is the provider key, and is the + handle written in DSL code (`a(V2E)`). +- `NeighborTable` / `NdArrayConnectivityField` — the **data**, unchanged, still + `Connectivity` implementations. + +This is the shape `FieldOffset` already has: it is *not* a `Connectivity` either, +and `premap` special-cases it at `nd_array_field.py:317-320`. So `a(V2E)` +continues to work by widening the same union — `Field.premap` and +`Field.__call__` are typed `Connectivity | fbuiltins.FieldOffset` +(`common.py:785, 791-794`) and become `Connectivity | type[NeighborConnectivity]` +in PR 4. `V2E` has **no instances**: `ConnectivityMeta.__call__` raises +`TypeError("… is a connectivity declaration and cannot be instantiated; bind a +table through the offset provider")`. + +Consequence: the proposal's sketch line +`class NeighborConnectivity(Connectivity[MultiDimensionIndex[Origin, Local], Codomain], ...)` +is **wrong and dropped**. `MultiDimensionIndex` remains the *domain index type* of +the `NeighborTable` (PR 8). **The knowledge-repo note needs this correction.** + +### 1.2 How `Local` reaches the base (proposal Open Q2) + +`requires-python = '>=3.12'`, so a PEP 696 default type parameter (3.13) is not +available. Resolution: **metaclass discovery**, base carrying a `ClassVar` +annotation, subclass declaring the nested class explicitly: + +```python +class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + Local: ClassVar[type[LocalDimensionIndex]] # annotation only, never assigned + +class V2E(NeighborConnectivity[V, E], max_neighbors=6): + class Local(LocalDimensionIndex): ... # explicit, required +``` + +Verified under `mypy --strict --python-version 3.12` and `pyright --pythonversion 3.12` +(`/tmp/probe_local.py`): + +| Variant | base declares | mypy | pyright | +| --- | --- | --- | --- | +| 1 | `Local: ClassVar[type[LocalDimensionIndex]]` | clean | clean | +| 2 | nothing | clean | clean | +| 3 | a real nested `class Local(LocalDimensionIndex)` | clean | **`reportIncompatibleVariableOverride`** | + +In all three the intended negative case (`Field[V, A.Local]` vs +`Field[V, B.Local]`) is correctly an error. Variant 3 is rejected. + +**Stated precisely — what variant 1 does and does not buy.** It does *not* make +`conn.Local` usable as a **type annotation** when `conn` is a generic +`type[NeighborConnectivity]`: both checkers reject that (mypy `name-defined`, +pyright `reportInvalidTypeForm`), and `T.Local` on a `TypeVar` is rejected too. +What variant 1 buys over variant 2 is only **value-level** access — +`reveal_type(conn.Local)` is `type[LocalDimensionIndex]` instead of an attribute +error — which is what library code in `common`, the backends and +`type_synthesizer` actually needs. Variant 1 is chosen for that, not for generic +annotations. Generic library code that must *name* a local dimension in a +signature uses `type[LocalDimensionIndex]`. + +This extends the proposal's probe P2: a **generated** `Local` is unusable as an +annotation, but a base `ClassVar` *annotation* plus an explicitly declared nested +class is fine. + +### 1.3 The IR keeps string tags; `resolve()` and `codegen_name()` are both required + +`AxisLiteral.value: str` stays (making it carry the class is a separate IR +change, deferred past this stack). It now holds the **qualified** tag, and that +has two consequences the first draft of this plan underestimated. + +**(a) `resolve(tag)` at every rebuild site**, memoized — `inference.py:464` calls +it once per `AxisLiteral` on the type-inference hot path: + +| Site | Purpose | +| --- | --- | +| `iterator/ir_utils/domain_utils.py` | `AxisLiteral` → `Dimension` | +| `iterator/ir_utils/misc.py` | `AxisLiteral` → `Dimension` | +| `iterator/type_system/inference.py:464` | `AxisLiteral` → `ts.DimensionType` | +| `codegens/gtfn/itir_to_gtfn_ir.py` (×2) | staggered-name sniffing → replaced in PR 2 by `Staggered[D]` | +| `dace/lowering/gtir_to_sdfg_lambda.py:1155` | synthesizes the local dim from the *offset* tag: `Dimension(offset, LOCAL)`. **Must be fixed in PR 2, not deferred** — see below | +| `dace/sdfg_args.py` | axis name → `Dimension` | +| `runners/roundtrip.py` | emits `gtx.Dimension(...)` as *source text* → becomes an import | +| ~~`common.flip_staggered` (×2)~~ | **not** a `resolve()` site: `Staggered[D]` replaces it with an interning subscript, §1.5 | + +`resolve` on a nested qualname was verified to work and round-trip +(`resolve("mymod.V2E.Local") is mymod.V2E.Local`), which matters because PR 4 keys +the provider on `V2E.Local.tag`. **One hazard to settle in PR 2**: a purely dotted +tag does not record *where* the module path ends and the qualname begins, so +`resolve` must try the longest importable prefix and walk the rest — O(depth) +import attempts, and in principle ambiguous if a module path and a class-attribute +chain collide. `pickle` avoids this by storing module and qualname *separately*. +Options: keep the pure dotted form (what the proposal asks for, ambiguity +tolerated and memoized away) or use an explicit separator such as +`"module:qualname"`. **Recommendation: keep the dotted form** — it is what makes +the tag "also a valid tag string for the IR", the collision requires a module and +an attribute chain to have the same spelling, and `resolve` can prefer the +*longest* importable prefix so a real module always wins. Record the residual in +the ADR. + +**(b) `codegen_name(tag)` — dots are illegal in every generated identifier.** +`eve`'s `SymbolName`/`SymbolRef` are constrained by +`_SYMBOL_NAME_RE = ^[a-zA-Z_]\w*$` (`eve/concepts.py:23,26,32`), so a qualified +tag reaching `Sym(id=...)` is a *validation error*, not a cosmetic problem. The +first draft mentioned mangling only in the abstract and put the roundtrip change +in a later PR; both were wrong. All of these are **PR 2**: + +| Site | What breaks without mangling | +| --- | --- | +| `codegens/gtfn/itir_to_gtfn_ir.py:170-195` | `TagDefinition(name=Sym(id=dim.value))` → `SymbolName` validation error | +| `codegens/gtfn/gtfn_module.py:97, 130-136` | `generated::{dim.value}_t`, plus `name.lower()` | +| `otf/binding/nanobind.py:197, 211` | C++ identifiers | +| `dace/lowering/gtir_to_sdfg_utils.py` `get_map_variable` | `i_{dim.value}_gtx_{kind}` → invalid DaCe symbol | +| `dace/sdfg_args.py:80` `_field_symbol` | invalid DaCe symbol | +| `dace/lowering/gtir_python_codegen.py:137-138` | `visit_AxisLiteral` returns the raw value | +| `runners/roundtrip.py:64, 177` | `AxisLiteral = as_fmt("{value}")`, and `{o.value} = gtx.Dimension(...)` emits `a.b.I = ...` → `SyntaxError` | + +**(d) The mangling scheme, corrected.** Earlier drafts said "injective (escape +existing `__` before replacing `.`)", i.e. `_ -> __` then `. -> _`. **That is not +injective**: `.` becomes a single `_`, so `".."` and `"_"` both map to `"__"`. +Exhaustively tested over the alphabet `{a, ., _}` up to length 6 +(`/tmp/probe_mangle.py`): **686 collisions in 1092 inputs.** Since a generated +identifier may only contain `[A-Za-z0-9_]`, `_` is the only available separator +and a *prefix escape* is required: + +```python +def codegen_name(tag: Tag) -> str: + return tag.replace("_", "_u").replace(".", "_d") + +def from_codegen_name(name: str) -> Tag: + return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) +``` + +Every `_` in the output is the first character of a two-character escape, so +decoding is unambiguous. Verified exhaustively over `{a, ., _, u, d}` up to +length 6 — **19530 inputs, 0 collisions, 0 round-trip failures**, every output a +valid identifier, including the adversarial `"_u"`, `"_d"` and `"a_ud.b"` +(`/tmp/probe_mangle2.py`). Cost: names grow (`mod.V2E.Local` → +`mod_dV2E_dLocal`), which is what gtfn's existing `TagDefinition.alias` mechanism +is for. + +**An inverse is needed too**, wherever generated names are parsed *back* into +dimensions: `dace/sdfg_args.py:25, 60-72` matches `gt_conn_(\S+)` and feeds the +result to `has_offset`. `codegen_name` must therefore be injective *and* have a +`from_codegen_name` partner (escape `__` → `____` before `.` → `__`). + +**A site that cannot be deferred: `gtir_to_sdfg_lambda.py:1155`.** It builds +`gtx_common.Dimension(offset, DimensionKind.LOCAL)` — a local dimension +synthesized from the **offset** tag, which in PR 2 is still a bare provider key +(`"V2E"`) that `resolve()` cannot import. Every DaCe unstructured shift passes +through it, so PR 2 is red on DaCe unless it is fixed there. The fix is local and +available: `conn_type` is already in scope (`:1134-1152`) and `:1135` already +asserts `conn_type.domain[1].kind == LOCAL`, so the line becomes +`offset_type = conn_type.domain[1]` (equivalently `conn_type.neighbor_dim`). +It is *necessary but not sufficient* for PR 1's `shift × tag≠localdim` DaCe cell: +that cell fails earlier, at `gtir_to_sdfg.py:842` +(`neighbor_table_types[dim.value]`, i.e. A4 on the connectivity *argument's* local +dim), before `:1155` is reached — and after the `:1155` fix, `:1371`/`:1455` would +reference `gt_conn_` while `:1104`/`:722` declare `gt_conn_`. So +**the DaCe shift cell stays in the skip matrix until PR 4**, where the +single-string choice makes both agree. (An earlier draft said PR 2; that would +leave PR 2 red on that cell.) + +**(c) The *offset* key is a second dotted name space, and it is mangled in PR 4, +not PR 2.** §1.3(b) covers *dimension* names only. When PR 4 makes the provider +key `cls.tag`, the **offset** string that flows through the IR +(`OffsetLiteral.value`, the provider key, `o` in the gtfn/DaCe connectivity +plumbing) becomes dotted too, and a different set of sites turns *it* into an +identifier. These are all **PR 4**: + +| Site | What breaks | +| --- | --- | +| `codegens/gtfn/itir_to_gtfn_ir.py:184` | `TagDefinition(name=Sym(id=offset_name))` → `SymbolName` regex | +| `codegens/gtfn/itir_to_gtfn_ir.py:490` | `SymRef(id=o)` for each connectivity → `SymbolRef` regex | +| `codegens/gtfn/codegen.py:147-148` | `visit_OffsetLiteral` emits `node.value` raw into C++ | +| `codegens/gtfn/gtfn_module.py:118, 132, 136` | `GENERATED_CONNECTIVITY_PARAM_PREFIX + name.lower()`, `generated::{name}_t` | +| `dace/sdfg_args.py:56` | `connectivity_identifier(name)` → `gt_conn_a.b.V2E`, an invalid SDFG array name | +| `dace/sdfg_args.py:60`, `dace/workflow/bindings.py:200, 286` | `is_connectivity_identifier` / `_parse_gt_connectivities` — the **inverse** direction, so `from_codegen_name` has *several* live consumers, not one | +| `dace/workflow/translation.py:61`, `dace/sdfg_callable.py:103`, `dace/program.py:156` | `connectivity_identifier(offset)` again, on the argument-binding path | +| `dace/lowering/gtir_to_sdfg_lambda.py:1104, 1371, 1455, 1727`, `gtir_to_sdfg.py:722` | the same identifier, consumed in the lowering | +| `runners/roundtrip.py:63` | `OffsetLiteral = as_fmt("{value}")` — emits the offset tag *raw as Python source*, into the program **body**; mangling `:176` alone still leaves `NameError: name 'tests' is not defined` | +| `dace/sdfg_args.py:83-84` | `_field_symbol`: `assert m[1] in offset_provider_type` — a *second* `from_codegen_name` consumer besides `:70` | +| `dace/lowering/gtir_to_sdfg_lambda.py:1892` | `visit_OffsetLiteral` → `SymbolExpr(node.value, INDEX_DTYPE)`, i.e. a dotted string used as a DaCe symbolic expression | +| `runners/roundtrip.py:152, 176` | collects offset-literal strings, then `f'{o} = offset("{o}")'` → `a.b.V2E = offset(...)` → `SyntaxError` | + +So `codegen_name` / `from_codegen_name` are introduced in PR 2 for dimensions and +**applied again in PR 4 for offsets**, at ~16 further sites. Two earlier claims +were wrong: that `from_codegen_name`'s only live consumer is in PR 2, and that the +DaCe surface is confined to `sdfg_args.py` and the lowering — the +argument-binding path (`workflow/translation.py`, `workflow/bindings.py`, +`sdfg_callable.py`, `program.py`) carries it too, in both directions. + +Because of (b) and (c), the review shortcut "diff PR 2 against #2844, the delta is +only identity" is **false**: #2844 needed none of this. Reviewers should expect a real +mangling layer on top of the identity delta. + +**`AxisLiteral.kind` becomes redundant** (the class carries it) — the `TODO` at +`iterator/ir.py:93`. Kept in PR 2, removed in PR 7, to keep PR 2's IR-expectation +churn to the `value` strings only. + +### 1.4 `LocalDimensionIndex` subclasses `DimensionIndex`; `DimensionBaseIndex` is dropped + +The proposal lists `DimensionBaseIndex` as a separate root with `DimensionIndex` +and `LocalDimensionIndex` as siblings. That does not survive contact with the +tree: `Dimension` is `type[DimensionIndex]`, eve validates a `type[X]` +annotation by `issubclass` (verified: a subclass passes, the base and an +unrelated class are both rejected), so sibling local dimensions would force +widening to `type[DimensionBaseIndex]` at `ts.DimensionType.dim`, +`ts.FieldType.dims`, `ConnectivityType.domain`, `Domain.__init__` and the `DimT` +/ `DimT_co` bounds — and would then accept local dimensions everywhere a primary +one is meant, which is the same looseness with extra ceremony. + +**Resolution**: `class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL)`. +`DimensionBaseIndex` is not introduced at all — one concept fewer, which is the +proposal's own stated goal. Where primary-only is required the check is +`dim.kind is not DimensionKind.LOCAL`, exactly as today. This also removes +#2844's deferral note ("a `DimensionBase` root, deferred until the requirements +of non-user-declarable dimensions are known") as a thing that needs resolving. + +**Verified**: all 38 sites in `src/` that discriminate a local dimension do so by +a **runtime `kind` check**, not by a static type distinction +(`transform_utils.py:65`, `type_deduction.py:460, 774`, +`custom_layout_allocators.py:171`, `past_to_itir.py:409`, `common.py:1168, 1336`, +`nd_array_field.py:972, 976`, `gtfn_module.py:91`, `embedded.py:922`, +`gtir_to_sdfg_types.py:76`, …). The tree already treats local dimensions as +`Dimension`s everywhere — `ConnectivityType.domain: tuple[Dimension, ...]` +includes the local one — so subclassing loses nothing it currently relies on, and +`Dims` (`tuple[Unpack[ShapeTs]]`, `common.py:57`) puts no bound on its members +either. + +**What subclassing does cost**, and the mitigation: every `DimensionIndex` +*bound* now statically admits a local dimension — +`NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]`, +`Staggered[D: DimensionIndex]` and `MultiDimensionIndex[D: DimensionIndex, *Ls]` +would all accept `V2E.Local` as their primary parameter. Each therefore gets a +runtime `kind is not DimensionKind.LOCAL` check in `__init_subclass__` / +`__class_getitem__`, and `LocalDimensionIndex.__init_subclass__` rejects an +explicit `kind=` other than `LOCAL`. This is the same runtime-check discipline +the tree already uses; the static gap is the price of the concept removed. + +**Deviation from the proposal; needs feeding back to the note.** + +### 1.5 `Staggered[D]` is required in PR 2, not PR 7 + +`flip_staggered` builds `Dimension(f"_Staggered{name}")` from a string +(`common.py:1452-1457`) and `is_staggered` tests `dim.value.startswith(prefix)` +(`:1447-1449`). #2844 routes both through the interning factory. With the +registry gone there is **no importable `_Staggered` type**, and a +dynamically created class would get the tag +`gt4py.next.common._Staggered`, so `is_staggered` is false and +`as_non_staggered` cannot recover the base dimension's module. Live dependents: +`test_staggered.py` (233 lines), `cases_utils.py:161` +(`KHalfDim = flip_staggered(KDim)`), gtfn `_add_staggered_aliases` +(`itir_to_gtfn_ir.py:203-215`), DaCe `get_map_variable` +(`gtir_to_sdfg_utils.py:52`), `type_synthesizer`, `test_common.py`, +`test_domain_utils.py`. + +So PR 2 is **not green** without `Staggered[D]`. It is Cartesian-only and does +not depend on the connectivity layer, so it moves into PR 2. + +**The obvious mechanism does not work.** A PEP 695 generic +`class Staggered[D: DimensionIndex](DimensionIndex)` makes `Staggered[KDim]` a +`typing._GenericAlias`, **not a class** (verified, `/tmp/probe_staggered.py`): + +``` +type(Staggered[KDim]) -> +isinstance(Staggered[KDim], type)-> False +issubclass(Staggered[KDim], ...) -> TypeError: issubclass() arg 1 must be a class +Staggered[KDim].tag -> '__main__.Staggered' # KDim is gone +``` + +So it fails eve's `type[DimensionIndex]` validation and its tag cannot name the +base dimension — it is not a `Dimension` at all. + +**The mechanism that does work** (verified, `/tmp/probe_staggered3.py`: runs +correctly and is **0 errors under both `mypy --strict` and pyright 1.1.414** on +3.12) is a metaclass `__getitem__` that *builds and interns a real class*, paired +with a `TYPE_CHECKING` declaration so checkers still see an ordinary generic: + +```python +class StaggeredMeta(DimensionMeta): + def __getitem__(cls, base: Dimension) -> Dimension: + if base not in _staggered_cache: + _staggered_cache[base] = StaggeredMeta( + f"Staggered[{base.__name__}]", + (cls,), # NOT (cls, base) -- see below + {"_tag": f"{cls.__module__}.{cls.__qualname__}[{base.tag}]", + "kind": base.kind, "base": base, "__slots__": ()}, + ) + return _staggered_cache[base] + +if TYPE_CHECKING: + class Staggered[D: DimensionIndex](DimensionIndex): + base: ClassVar[Dimension] +else: + class Staggered(DimensionIndex, metaclass=StaggeredMeta): + __slots__ = () + base: ClassVar[Dimension] +``` + +Verified properties of `Staggered[KDim]`: it *is* a class; +`tag == "gt4py.next.common.Staggered[]"`; `kind` is +inherited from the base; `issubclass(_, DimensionIndex)` and +`issubclass(_, Staggered)` hold; it is instantiable as an index; and +`Staggered[KDim] is Staggered[KDim]`, so identity is stable. `Staggered[KDim]` in +an annotation and inside `Field[Dims[Staggered[KDim]], float]` are both accepted +by both checkers. + +- **Bases are `(cls,)`, not `(cls, base)`.** Inheriting from the base dimension + would make `issubclass(Staggered[KDim], KDim)` true, i.e. `KHalfDim` would be + accepted everywhere `KDim` is required. It is a *different* dimension; only + `kind` is inherited, copied explicitly into the namespace. +- `is_staggered(dim)` becomes **`"base" in dim.__dict__`**, not + `issubclass(dim, Staggered)`, and `as_non_staggered(dim)` becomes `dim.base`. + Two runtime facts force this: `issubclass(Staggered, Staggered)` is true for + the bare base, which has no `base`; and `Staggered[KDim]` is *subclassable* + (`class KHalf2(Staggered[KDim])` yields a second, un-interned staggered-K type + with tag `.KHalf2`). `Staggered.__init_subclass__` therefore rejects + any subclass the metaclass did not create, so the interned form is the only + one. Still structural — no string sniffing. +- **The guards were verified, including the escape routes** + (`/tmp/probe_staggered_guards.py`). All four are blocked: + `class KHalf2(Staggered[KDim])`, `class X(Staggered)`, a direct + `StaggeredMeta("Y", (Staggered,), {})`, and the double subscript + `Staggered[KDim][KDim]`. The `copyreg` fallback round-trips the bare + `Staggered` by reference, the parametrized class with identity preserved, and + instances. Implementation note: gate `__init_subclass__` on a **namespace + marker** the metaclass sets (`"_tag" in cls.__dict__`), not on a module-level + "currently building" flag — the flag works but is not thread-safe, and + compilation runs in worker processes and threads. Three further honest limits: + the guard defends against **accidental** subclassing only — a deliberate + `StaggeredMeta("Forged", (Staggered,), {...marker})` or `types.new_class` can + still forge a same-`tag`, non-identical type (as it can for any class); + `Staggered[Staggered[KDim]]` must be rejected explicitly by testing + `"base" in base.__dict__` in `__getitem__`, or it nests and pickles happily; and + a hand-built `copyreg` payload such as `(_make_staggered, (int,))` should raise + a `TypeError` naming the offending base rather than an `AttributeError`. +- `resolve` gains the `[]` grammar: it parses the brackets and + evaluates `Staggered[resolve(inner)]`, which hits the same intern cache, so a + staggered dimension round-trips through the IR to the *same* class object. +- **A narrow `copyreg` is required after all.** `Staggered[KDim]`'s + `__qualname__` is `Staggered[KDim]`, which `pickle.save_global` cannot look up: + `PicklingError: Can't pickle : attribute lookup + Staggered[KDim] on … failed` (verified, `/tmp/probe_staggered_pickle.py`). A + `copyreg.pickle(StaggeredMeta, lambda cls: (_make_staggered, (cls.base,)))` + fixes it *and preserves identity*, because the reconstructor goes back through + the intern cache. **But it must guard the bare base**: `type(Staggered) is + StaggeredMeta` too, so a reducer that unconditionally reads `cls.base` fails on + `Staggered` itself with `AttributeError: type object 'Staggered' has no + attribute 'base'` (verified — an earlier draft of this section claimed the + registration "captures only parametrized dimensions", which is false). The + reducer therefore falls back to by-reference pickling when + `"base" not in cls.__dict__`. It never captures a plain dimension + (`type(KDim) is DimensionMeta`). This is materially narrower than + #2844's blanket registration on `DimensionMeta` — a parametrized type needs a + reconstructor for the same reason `typing` aliases do — but §0's "`copyreg` + dropped" row is only true of the blanket form, and the ADR must say so. +- **Two honest costs.** (i) `_staggered_cache` is a cache, and the proposal's + headline is that the *name-keyed* registry goes away. The difference is real but must be + stated: it is keyed by a *dimension class*, is internal, and is memoization of + a type constructor (as `typing`'s own subscription cache is), not interning of + user-authored name strings — nothing resolves a user string through it. + (ii) the `TYPE_CHECKING` split means the static and runtime definitions can + drift; a unit test must assert the runtime facts the static form does not + express (real class, `issubclass` against `Staggered` but *not* against the + base, tag shape, interning). +- Supersedes ADR 0026, recorded in the PR-2 ADR. + +## 2. The PR stack + +Branches follow the repo's stacked convention, `connectivities-as-types--`, +each based on its predecessor, all targeting `main`. PR titles are Conventional +Commits (squash-merge lands the title). + +--- + +### PR 1 — `fix[next]: lower unstructured shifts with the offset's own tag` + +**Independent of the rest of the stack; lands first, on its own merit.** + +`foast_to_gtir._visit_shift` emits the **Python variable name** as the IR shift +tag (`foast_to_gtir.py:305` `offset_name.id`, `:331` `str(offset_name)`), because +`ts.OffsetType` does not carry the tag. Embedded execution keys on +`FieldOffset.value`. So the same program needs a *different* provider key +depending on the backend — confirmed by running it on v1.2.2: + +``` +MyOff = FieldOffset("TAGNAME", ...) +embedded: {"TAGNAME": conn} OK ; {"MyOff": conn} -> KeyError 'TAGNAME' +roundtrip: {"MyOff": conn} OK ; {"TAGNAME": conn} -> KeyError 'MyOff' +``` + +**Change** + +- `ts.OffsetType` gains **`tag: Optional[Tag] = None`** — *not* a required field. + `type_deduction.py:709` builds `ts.OffsetType(source=conn.codomain, + target=(conn.domain_dim,))` from `IDim + 1`, a `CartesianConnectivity` that has + no tag at all; making `tag` required breaks it. +- `FieldOffset.__gt_type__` fills it (`fbuiltins.py:485`). +- `type_deduction.py:464`, which rebuilds an `OffsetType` when `Off[1]` drops the + local dimension, must **propagate** the tag. +- `foast_to_gtir._visit_shift`: the `Subscript` branch and the bare `Name` branch + use `arg.type.tag`, asserting non-`None` (both are unstructured paths, where a + tag always exists). + +**Tests.** `tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py` +today covers exactly `a(Off[1])` on `GTFN_CPU`. Extend to +{shift, `neighbor_sum`} × {embedded, roundtrip, gtfn, dace} × {tag≠varname, +tag≠local-dim-name}. + +**The matrix is not uniform, and a blanket `xfail` will not do.** `xfail_strict = true` +(`pyproject.toml:323`), and measured behaviour on v1.2.2 is: + +| case | embedded | roundtrip | gtfn | dace | +| --- | --- | --- | --- | --- | +| shift, tag≠varname | pass | pass | pass | pass | +| shift, tag≠localdim | pass | pass | pass | **fail** `KeyError` (`gtir_to_sdfg_lambda.py:1155` synthesizes the local dim from the tag) | +| `neighbor_sum`, tag≠localdim | **fail** | **pass** | **fail** | **fail** | + +So the first draft's acceptance criterion ("shift cells pass on all four +backends") is unreachable before the backend work, and a strict blanket `xfail` +would XPASS on roundtrip. **Fix**: add a per-backend skip matrix entry in +`tests/next_tests/definitions.py` (a new `USES_*` marker) covering exactly the +failing cells, roundtrip excluded. **They are removed in two steps**: the +`shift × tag≠localdim` DaCe cell and the three `neighbor_sum × tag≠localdim` cells +all in **PR 4**, where the single-string choice makes A3/A4 vacuous — *not* in +PR 5, and *not* the DaCe cell in PR 2 (the `:1155` fix there is necessary but not +sufficient; see §1.3(a)). The gtfn `neighbor_sum` +failure is now confirmed **by running it**; the proposal had it only "by +reading". + +**No CHANGELOG entry.** Verified against the history: `CHANGELOG.md` is touched +*only* by release PRs (`git log -- CHANGELOG.md` is release commits exclusively, +and nothing between `b3c53fa7e` and `upstream/main` touches it). The behaviour +change — which key a compiled backend requires when tag ≠ variable name — belongs +in the PR description, and reaches the changelog when the release PR is cut. Two +earlier drafts of this plan said otherwise, including for PR 6's breaking change. + +**ICON4Py is unaffected by PR 1**: all 16 `FieldOffset` variable names equal +their tags (`model/common/src/icon4py/model/common/dimension.py:33-48`). + +**Acceptance**: `nox -s test_next` green; every cell in the matrix either passes +or is covered by the documented skip matrix. + +--- + +### PR 2 — `feat[next]: a concrete Dimension is a class, identified by its qualified name` + +The #2844 core with the identity divergences of §0, **plus** the mangling layer +of §1.3(b) and `Staggered[D]` of §1.5 — both of which #2844 did not need and +without which this PR cannot be green. Large and largely mechanical. + +**`src/gt4py/next/common.py`** + +```python +class DimensionMeta(type): + kind: DimensionKind + __hash__ = type.__hash__ # §1.0 — mandatory, not optional + @property + def tag(cls) -> Tag: ... # f"{cls.__module__}.{cls.__qualname__}" + # operators as in #2844; __eq__/__ne__ keep only the IntegralScalar overload + # (I == 5 -> Domain); the dim-vs-dim branch is `is`. + +class DimensionIndex(metaclass=DimensionMeta): + __slots__ = ("value",) + kind: ClassVar[DimensionKind] = DimensionKind.HORIZONTAL + def __init_subclass__(cls, /, kind=None, **kw): ... + +# Staggered: an interning metaclass + TYPE_CHECKING split, NOT a PEP 695 +# generic -- see §1.5, where the generic form is shown to be unworkable. + +def resolve(tag: Tag) -> Dimension: ... # memoized; [] grammar +def codegen_name(tag: Tag) -> str: ... # "_" -> "_u", "." -> "_d" (§1.3(d)) +def from_codegen_name(name: str) -> Tag: ... # the inverse, `_([ud])` -> `_` / `.` + +type Dimension = type[DimensionIndex] +``` + +- `tag` is a metaclass **property**, so it cannot drift from the type. This makes + a class-body `tag = "C2E"` a **silent no-op** (verified: `C2EDim.tag` stays + `"__main__.C2EDim"` even with `tag = "C2E"` in the body) — and that pattern is + exactly what ICON4Py and #2845 use to rename. `__init_subclass__` therefore + **raises** on `"tag" in cls.__dict__`, naming the class and pointing at the + rename path. +- `__init_subclass__` also rejects `"" in cls.__qualname__`. Neither + necessary nor sufficient (`type("Dyn", ...)` in a function passes; a `del`'d + class passes) — the authoritative check stays pickle's own `save_global`. +- `resolve` raises a `ValueError` naming the tag and the failing import, per + CODING_GUIDELINES. + +**Removals**: `NamedIndex`; `_DimA`..`_AnyDim` and the dimension half of +`mypy_plugin.py`; `_DIMENSION_REGISTRY`; `copyreg`; the `fingerprinting.py` +deconstructor; `_STAGGERED_PREFIX` and its string sniffing. + +**Migration**. Every `Dimension("X")` becomes `class X(DimensionIndex): ...` at +module level. Verified counts: 333 `Dimension("` declarations in `tests/`, of +which **133 are function-local across 15 files** and must move to module level; +a dimension *named* `"I"` is declared 46 times across **17** files (the first +draft said 45 files — that was the proposal's *`IDim` file* count, a different +number). Docs, workshop notebooks and `examples/` are included because +`test_examples` executes them; notebook *code* cells only, stored outputs +untouched (they hold recorded tracebacks that must keep naming the symbols that +produced them). + +**IR expectation churn**: 36 `AxisLiteral` and 37 `OffsetLiteral` occurrences in +`tests/`, most already computed from `dim.value`. `test_pretty_roundtrip.py` and +the gtfn/DaCe snapshot tests hold the hardcoded names. + +**Do the sweep with a codemod script, not agent fan-out.** A previous attempt at +agent fan-out on a large mechanical rewrite in this repo died mid-file on the +rate limit and left the tree inconsistent; a script did all 57 files uniformly. + +**ADR 0028** (the directory ends at 0027): nominal identity; the module-level +declaration requirement; `Staggered[D]` superseding ADR 0026; that cache +fingerprints now shift when a declaration moves module (a consequence for ADR +0023, not a reversal); that `resolve()` imports modules named in the IR, which is +the same trust level as `pickle` loading a class by reference. + +**Documented limitation**: interactive `__main__` (REPL, notebooks, `python -c`) +cannot be resolved. `spawn` compile workers re-execute the main *script* as +`__mp_main__`, so file-based `__main__` resolves provided the script has the +`if __name__ == "__main__":` guard the pool already requires. + +**ICON4Py migration script is a PR-2 deliverable, not PR 6.** All 15 local +dimensions and `KDim`/`EdgeDim`/`CellDim`/`VertexDim` have variable name ≠ tag +(`EdgeDim = Dimension("Edge")`), so PR 2 changes every generated symbol and every +cache key downstream. + +**Acceptance**: `nox -s test_next` on **3.12, 3.13 and 3.14** (the `typing` +subscription cache behaves differently per interpreter and this change moves +exactly that behaviour), then `test_eve`, `test_storage`, `test_cartesian`, +`test_examples`; `uv run mypy src/`; `uv run pyright`; `uv run tach check`; +`uv run pre-commit run -a`. One at a time, pytest capped at `-n 4`. + +--- + +### PR 3 — `feat[next]: NeighborConnectivity declarations and local dimensions that know their owner` + +**Purely additive**: new concepts next to `FieldOffset`, nothing removed, no +behaviour change, no test churn. + +```python +class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): # §1.4 + owner: ClassVar[type[NeighborConnectivity] | None] = None + max_neighbors: ClassVar[int | None] = None + min_neighbors: ClassVar[int | None] = None + def __init_subclass__(cls, *, size: int | None = None, **kw): ... + +class ConnectivityMeta(type): # §1.0 for __hash__ and __getitem__ + @property + def tag(cls) -> Tag: ... + def __call__(cls, *a, **kw) -> NoReturn: ... + +class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + Local: ClassVar[type[LocalDimensionIndex]] + def __init_subclass__(cls, *, max_neighbors=None, min_neighbors=None, **kw): ... +``` + +- `__init_subclass__` asserts `"Local" in cls.__dict__` and that it subclasses + `LocalDimensionIndex`, then sets `Local.owner = cls` and copies the counts. A + missing `Local` is a `TypeError` at class creation naming the class. +- **Owner-less** locals: `class LsqUnk(LocalDimensionIndex, size=3)` — `owner is + None`, `min == max == size`, never in the provider. ICON4Py's `LsqUnkDim` and + `RBFDimension` need this: they index no table but need sparse storage and + layout. +- Counts are **optional class keywords**, not type parameters (Python has no + integer type parameters and nothing static needs the count). Declared ⇒ a + constraint the table must satisfy. Undeclared ⇒ completed at bind time, from + the table in the JIT flow or from the `NeighborConnectivityType` already passed + through `connectivities=` (`ffront/decorator.py:188-208`) in the AOT flow. + Not static-only because `fvm_nabla_setup.py:99` sizes `V2E` from the atlas + mesh, and ICON skip-value presence is configuration-dependent (`icon.py:130` — + pentagons have skip values on the icosahedron, not the torus). +- **Bind-time validation**, one function replacing constraints A6–A8: shape + `(n, max_neighbors)`, integral dtype, skip values present iff + `min_neighbors < max_neighbors`, `domain[0] is Origin`, `codomain is Codomain`. + +**No provider bridge.** The first draft proposed a dual-keyed +(`Tag | type[NeighborConnectivity]`) provider here. Dropped: the provider is +accessed **directly, not through `get_offset`, at 19 sites in 12 `src/` files** +(`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, +`compiled_program.py`, `pass_manager.py`, …) despite the note at +`common.py:1174`, so a bridge would be both invasive and — since nothing would +exercise class keys — untested. Class-keyed providers land in one place, PR 4. + +**Typing tests**: `typing_probe.py` / `probe_local.py` / `probe_meta_getitem3.py` +become real coverage — `typing_tests/test_next.yaml` cases for +`Field[Dims[V, V2E.Local], float]`, a `TypeVar` bound to `LocalDimensionIndex`, +the negative cross-connectivity case, and the `NC[V, E]`-in-bases case of §1.0; +plus the pyright variants. + +**Acceptance**: full suite green with no behaviour change; new unit tests for +declaration errors, owner wiring, owner-less locals and bind-time validation. + +--- + +### PR 4 — `feat[next]: declare connectivities as classes; FieldOffset derived from them` + +**The ordering fix.** Revision 2 put class-keyed providers here and the +declaration migration in PR 6. That cannot be green: `Local.owner` only exists if +the user declared a `NeighborConnectivity` class, and a `FieldOffset` written the +old way (`FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim))`) has no class +to point at — so neither the backend work nor a class-keyed provider has anything +to resolve. The declaration migration must come **first**, and the provider key +must stay a string until the backends are through. + +- The **unstructured** `FieldOffset` is **derived**, not authored: + `FieldOffset.from_connectivity(V2E)` (or `V2E.__gt_offset__()`), which fills + `source = Codomain`, `target = (Origin, V2E.Local)` and — critically — + **`value = V2E.Local.tag`, the *local dimension's* tag, not `V2E.tag`.** +- **Why the local dimension's tag and not the connectivity's.** An earlier draft + used `V2E.tag` and claimed PR 4 was green. It is not: A3 and A4 key the provider + on the **local dimension's** name, at `nd_array_field.py:983` + (`get_offset(provider, axis.value)`, whose in-tree comment is literally + `# assumes offset and local dimension have same name`), `unroll_reduce.py:47` + (`arg.type.offset_type.value`), `gtfn_module.py:95`, `gtir_to_sdfg.py:581, 842`, + `iterator/embedded.py:954, 1519`, and + `gtir_to_sdfg_lambda.py:1371, 1455` (`connectivity_identifier(offset_type.value)`). + Today the `V2EDim = Dimension("V2E")` convention makes that string equal to the + offset tag; PR 4 deletes the convention tree-wide, while the `owner` lookup that + replaces it is PR 5. With `value = V2E.tag` every reduction and every sparse-field + argument would break on embedded, gtfn and DaCe simultaneously — the round-1 + matrix row (`neighbor_sum`, tag≠localdim: embedded/gtfn/DaCe fail) would become + the tree's universal state. + Choosing `V2E.Local.tag` instead makes **all four** of A1, A3, A4 and A5 vacuous + at once, because there is then exactly *one* string and the class produces it. + It also makes the #1789 branch at `itir_to_gtfn_ir.py:181-190` + (`if offset_name != connectivity_type.neighbor_dim.value`) dead already in PR 4. + This is preferable to the alternatives — fusing PR 5 into PR 4, or a transient + `owner`-based fallback inside `get_offset` — because it needs no scaffolding: + the string is simply picked correctly, and PR 5 then removes the dependence on a + string at all. +- **The Cartesian `FieldOffset` constructor stays in PR 4.** An earlier draft said + "every declaration becomes a class", which is wrong: `Ioff`, `Koff` and + `EdgeOffset` (`cases_utils.py:163-169`, e.g. + `Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,))`) and ICON4Py's + `Koff`/`KHalfOff` (`dimension.py:47-48`) are single-target and have no + `NeighborConnectivity` to derive from, and their only remaining consumer — + `as_offset` — does not change until PR 6. Restricting the public constructor to + the single-target form keeps PR 4 green; it disappears with `as_offset` in + PR 6. +- **Providers stay keyed on `Tag`**, now `V2E.Local.tag`. Nothing about the key + *mechanism* changes yet, so the 19 direct-access sites are untouched. A1, A3, A4 + and A5 are all dead at this point — there is one string, and the class produces + it. +- `ts.OffsetType` → `ConnectivityType`, produced by the class + (`type_specifications.py:74` TODO). `type_info.py:637, 858` gate `a(V2E)` + deduction on `ts.OffsetType` and follow. `Field.premap` / `Field.__call__` + unions widen to `Connectivity | type[NeighborConnectivity]` (`common.py:785, + 791-794`, `1313-1316`). +- **Test-tree migration lands here**: `toy_connectivity.py`, `cases_utils.py`, + `fvm_nabla_setup.py` are the fixture modules everything imports; 37 + `FieldOffset` sites, 42 `DimensionKind.LOCAL` sites. String provider keys keep + working because they are `cls.tag` — but the tags are now *qualified*, so the + 111 ITIR-level string-offset occurrences in 15 files (`im.shift("V2E")`, + `neighbors("…")`, `OffsetLiteral(value="…")`, string-keyed providers) are + rewritten to `V2E.Local.tag` here rather than in PR 6. +- **ICON4Py**: this is the release-visible declaration change. The migration + script written in PR 2 is extended. + +--- + +### PR 5 — `refactor[next]: backends resolve connectivities through the local dimension's owner` + +Where A3, A4 and A5 dissolve and PR 1's skip-matrix entries are removed. Green +while providers are still string-keyed, because a backend goes +`local_dim.owner` → `owner.tag` → the existing lookup: the *identity* question is +answered by the owner pointer, and the key is still a string. + +**The owner-less case must be handled, not assumed away.** At PR 5 `_CONST_DIM` +is still a plain `DimensionIndex(kind=LOCAL)` (it becomes `ConstList` only in +PR 7), and `LsqUnk`-style local axes have `owner is None` by design. Every +converted site reads `getattr(dim, "owner", None)` and falls through when it is +`None`. Two sites already guard by accident — DaCe compares against `_CONST_DIM` +first (`gtir_to_sdfg_lambda.py:1314`) and `unroll_reduce` filters +`offset_type is None` — but `gtfn_module.py:91-98` and `nd_array_field.py:981` +have **no** guard, and ICON4Py never exercises the case +(`test_icon.py:220`), so the gap would not show up downstream. + +`unroll_reduce.py:47` (reads `arg.type.offset_type`, which is the local +`Dimension` — now a `LocalDimensionIndex` carrying `owner`, which is exactly the +back-pointer it lacked), `gtfn_module.py:95, 118, 132`, `itir_to_gtfn_ir.py` +(including the `#1789` `offset_name != neighbor_dim.value` branch at `:181-190`, +which becomes dead and goes), `gtir_to_sdfg.py:581`, +`gtir_to_sdfg_lambda.py:766-770, 1155` (the `Dimension(offset, LOCAL)` synthesis +goes), `nd_array_field.py:981-985` (and its +`# assumes offset and local dimension have same name` comment), +`iterator/embedded.py`, and `runners/roundtrip.py`. + +--- + +### PR 6 — `feat[next]!: class-keyed offset providers; remove FieldOffset and the string offset API` + +The only breaking PR, and now the only one that touches the provider key. + +- `OffsetProvider*` become + `Mapping[type[NeighborConnectivity], NeighborTable]`; `get_offset` keys on the + class. **Note the `.owner` hop**: because PR 4 made the IR offset tag the *local + dimension's* tag, `resolve(OffsetLiteral.value)` yields a `LocalDimensionIndex`, + not the connectivity — so the class-key lookup is + `resolve(tag).owner`. (The alternative is to switch the IR tag to `V2E.tag` in + this PR; the `.owner` hop is cheaper and keeps the IR stable.) The **19 direct-access sites in 12 `src/` files** + (`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, + `compiled_program.py`, `pass_manager.py`, …) are converted here, and + `common.py:1174`'s "all accesses should go through `get_offset`" either becomes + true or the note goes. `hash_offset_provider_items_by_id` and the + `fingerprinting` dict handling already tolerate class keys once §1.0's + `__hash__` is in place. +- **Removals**: `FieldOffset` entirely (both forms); `runtime.Offset` as its base + (`fbuiltins.py:467-470` TODO); `iterator/runtime.offset("...")` (12 sites in 6 + files, plus `tracing.py:161-162`); the `V2EDim`-next-to-`V2E` convention; + `embedded/context.py` string plumbing; the `gt4py.next.__init__` exports at + `:47, 140`. +- **`as_offset` changes in the same PR.** It is why the Cartesian `FieldOffset` + form cannot go alone: `ffront/experimental.py:17` + + `type_deduction.py:956-967` require one. New signature + `as_offset(KDim, field)`. Used in 5 test modules, the `Ioff`/`Koff`/`EdgeOffset` + fixtures at `cases_utils.py:163-169`, and **40 non-test call sites in + ICON4Py**. +- `transform_utils.py:50-77` and `past_to_itir.py:77` deduce grid type from the + provider and follow. +- **Accepted double churn**: the ~26 `offset_provider={...}` literals are rewritten + twice — `{V2E.Local.tag: t}` in PR 4, `{V2E: t}` here. The alternative is fusing PR 4 + and PR 6, which loses the green boundary. The ITIR string sites do *not* churn + twice: `im.shift(V2E.Local.tag)` written in PR 4 stays correct. + +**ADR 0029**: the connectivities-as-types record — `FieldOffset` removed, the +class-keyed provider, superseding the `FieldOffset` part of ADR 0019. + +**Breaking-change communication**: the PR title carries the Conventional Commits +`!` marker and the ADR records the removal; the changelog entry is written by the +release PR, not here (see PR 1). No deprecation window — an explicit decision: +ICON4Py's provider keys are bare names, so no import-based shim could have +resolved them. + +**Acceptance**: full suite; `test_fvm_nabla` and `test_icon_like_scan` are the +integration canaries. A before/after run of one gtfn and one DaCe program +checking generated-code equivalence modulo names. + +### PR 7 — `refactor[next]: ConstList replaces _CONST_DIM; AxisLiteral drops kind` + +- `_CONST_DIM` (`iterator/embedded.py:220` and + `dace/lowering/gtir_to_sdfg_lambda.py` — two separate declarations, each + internally consistent; 14 references in total) becomes the owner-less + `ConstList(LocalDimensionIndex, size=1)`, generalizing the magic name from size + 1 to size *n*. +- `AxisLiteral.kind` removed (`iterator/ir.py:93` TODO), now that every tag + resolves to a class carrying its kind. + +--- + +### PR 8 — `refactor[next]: MultiDimensionIndex and typed embedded positions` + +- `MultiDimensionIndex[D: DimensionIndex, *Ls]` as the index type of a sparse + position and the domain index of a `NeighborTable`. `*Ls` is unconstrained + because `TypeVarTuple` cannot carry a bound; `__init_subclass__` checks at + runtime what the checker cannot. +- `iterator/embedded.py` positions keyed by dimension types instead of name + strings (`embedded.py:574-576`, `597-616`, `941-950`); `SparseTag` removed. + Constraint A9 dissolves. +- Nothing else depends on this; it is last for that reason. + +--- + +## 3. Constraint ledger + +| # | Constraint | Retired by | +| --- | --- | --- | +| A1 | `FieldOffset.value` == provider key | PR 4 (one declaration produces both) | +| — | *all four string-equality constraints below become vacuous in PR 4*, because the class emits a single string (`V2E.Local.tag`); PR 5 removes the dependence on a string at all | PR 4 / PR 5 | +| A2 | Python variable name == provider key | **PR 1** | +| A3 | local dim name == provider key (reductions) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | +| A4 | local dim name == provider key (sparse args) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | +| A5 | `FieldOffset.value` == local dim name | PR 4 (one declaration) | +| A6 | `target[-1]` == connectivity `neighbor_dim` | PR 3 (bind-time check) | +| A7 | `FieldOffset.source` == `codomain` | PR 3 (bind-time check) | +| A8 | `target[0]` == `domain[0]` | PR 3 (bind-time check) | +| A9 | dim name is the iterator-position dict key | PR 8 | +| A10 | dim name round-trips through `AxisLiteral` | structural; PR 2 qualifies it, PR 7 drops `kind` | +| F5/F8/F9 | codegen name formats | PR 2 (`codegen_name` + inverse) | +| S6 | `as_offset` needs a Cartesian `FieldOffset` | PR 6 | +| — | string-keyed provider | PR 6 | + +## 4. Risks + +1. **PR 2 size, and it is not purely mechanical.** ~150 files of migration *plus* + a name-mangling layer and `Staggered[D]`. It cannot be reviewed as "the + #2844 diff plus identity". Mitigation: land the mangling layer and + `Staggered[D]` as reviewable commits *within* the PR, ordered before the + sweep, so the mechanical part is a separate commit. +2. **Function-local dummy dimensions.** 133 declarations in 15 test files must + move to module level. Under nominal identity, two same-named locals that were + silently the same dimension become distinct — each resulting failure is a + real finding, not churn. +3. **`typing` subscription caching.** Under `(name, kind)` equality, + `Field[Dims[I]] is Field[Dims[I2]]` aliases for two distinct same-named + classes — a known residual of the #2844 design, and an argument *for* this + stack. The cache behaves differently per interpreter: verify on 3.12, 3.13 + and 3.14 **via nox**, not `uv run pytest`. +4. **Fingerprint/cache invalidation.** Moving a declaration between modules now + invalidates compiled artifacts. Intended; CHANGELOG + ADR line. +5. **Naming not yet converged** with `havogt/dependent-local-dimensions` + (`Origin`/`Codomain` vs `source_dim`/`neighbor_dim`; `Local` vs `Dim`; + `min_neighbors` vs `has_skip_values`). PR 3 fixes public names. Converge + before PR 3 is *opened*. +6. **`V2E.Local` vs `Local[V2E]`.** The chain proposals' encodings subscript + `Local`, and a `TypeVar` cannot be subscripted for a nested attribute. Their + semantics are unaffected; their static encoding needs rewriting to `C.Local` + plus a protocol for the generic hop-stack case. Knowledge-repo concern, not a + gt4py blocker. +7. **CSCS GPU CI is flaky and opaque.** All jobs failing at the same second means + infrastructure; `cscs-ci run default` as a PR comment reruns it. #2844's CI is + green except that job. +8. **Two mangling passes, two PRs.** `codegen_name` is applied to dimension names + in PR 2 and to offset names in PR 4, at ~19 sites total, several of which + (`Sym`/`SymRef` construction, DaCe array names) fail *loudly* and several of + which (C++ emission, `name.lower()`) fail only in the generated artifact. + Both PRs need a test that a qualified tag survives a real gtfn and a real + DaCe compile, not just lowering. +9. **`resolve()` on the inference hot path** (`inference.py:464`, once per + `AxisLiteral`). Must be memoized from the start, and the memo must be keyed + so a reloaded module does not return a stale class. + +## 4b. Work the earlier drafts did not mention + +- **Public exports.** `gt4py.next.__init__` must export `NeighborConnectivity`, + `LocalDimensionIndex`, `Staggered` and `resolve` (PR 2 for the dimension half, + PR 3 for the connectivity half), and drop `FieldOffset` / `offset` at `:47, 140` + in PR 6. +- **`type_translation.from_value(V2E)`** works only because the + `hasattr(value, "__gt_type__")` branch at `type_translation.py:328` is tested + *before* the `DimensionMeta` branch. That ordering is load-bearing under this + design and currently untested — PR 3 adds a unit test pinning it. +- **`pyright` is not yet a dependency.** `uv run pyright` appears throughout §5 + but pyright is absent from `pyproject.toml` on `main`; #2845 is what adds it. + PR 2 must explicitly fold in #2845's `typing_exports` / pyright dependency-group + change, or §5's pyright step is not runnable. +- **`test_examples` belongs to PR 4 too.** §5 lists it for PR 2 and PR 6; the + docs and notebooks use `FieldOffset`, so PR 4's declaration migration touches + them and must run it. + +## 5. Verification + +Per PR, in this order, **one at a time** on the shared machine, pytest capped at +`-n 4`: + +``` +uv run pre-commit run -a # ruff, mypy, tach, license headers +uv run pyright # static checks the mypy plugin no longer fakes +uv run nox -s "test_next-3.12(...)" # then 3.13, 3.14 for PR 2 +uv run nox -s test_eve test_storage test_cartesian test_examples # PR 2, PR 6 +``` + +Test-first where behaviour changes, per AGENTS.md: PR 1's regression matrix, PR +3's declaration-error and bind-validation units, and PR 4's provider-key tests +are written before the implementation they cover. + +## 6. Feedback owed to the knowledge-repo note + +- `NeighborConnectivity` is not a `Connectivity`; the sketch's base line is wrong + (§1.1) — this closes Open Q6. +- `DimensionBaseIndex` should be dropped; `LocalDimensionIndex` subclasses + `DimensionIndex` (§1.4). +- Open Q2 is closed: metaclass discovery, base `ClassVar` *annotation*, explicit + nested class — with the precise limit of what that buys statically (§1.2). +- `Staggered[D]` is not a late step; it is a precondition for removing the + name-keyed registry (§1.5) — and it cannot be a PEP 695 generic, needs an + identity-keyed intern cache, and needs a narrow `copyreg`. The note's claim + that `copyreg` disappears entirely is therefore too strong. +- The note's "five name spaces" analysis should record that under qualified tags + there are **two** dotted name spaces reaching codegen — dimension tags and + offset tags — each needing its own mangling pass (§1.3(b) and (c)). +- The note's §Staging step 2 ("`NeighborConnectivity` … object-keyed provider; + `FieldOffset` and string keys removed outright") bundles three changes that + must be separated to stay green: the declaration migration has to precede the + backend work (because `Local.owner` only exists once classes are declared), and + the provider *key* has to stay a string until the backends resolve through the + owner. See PR 4/5/6. +- The staging in the note's §Staging (steps 0–8) is superseded by §2 here; in + particular step 4 ("backends, one file at a time") cannot follow step 2, since + the provider key change and the backend lookups are separable but the mangling + layer is needed at the *dimension* step. + +## 7. Open, non-blocking + +- Naming convergence (risk 5). +- Whether a same-`__name__` collision warning is useful or noise (note Q5). +- Whether interactive `__main__` should be detected with a fallback to in-process + compilation and a warning, or merely documented (note Q3). Plan assumes + documented. +- `AxisLiteral.dim` instead of `AxisLiteral.value` — a follow-up after PR 8. diff --git a/docs/development/next/fieldoffset-tag-constraints.md b/docs/development/next/fieldoffset-tag-constraints.md new file mode 100644 index 0000000000..45122c5a49 --- /dev/null +++ b/docs/development/next/fieldoffset-tag-constraints.md @@ -0,0 +1,551 @@ +# Tag and name constraints for `FieldOffset`, `Dimension` and offset providers + +Analysis of the *string-identity* assumptions that connect `FieldOffset` tags, +`Dimension` names and `offset_provider` keys in `gt4py.next`, across the +embedded, IR and backend contexts. + +Status: descriptive — this documents the implementation as it is, it does not +propose a change. Line references are against `main` at `b3c53fa7e` +(v1.2.2, 2026-09-03). + +Related: [ADR 0019 Connectivities](../ADRs/next/0019-Connectivities.md), +[ADR 0026 Staggered Dimensions](../ADRs/next/0026-Staggered_Dimensions.md). + +## 0. The five name spaces + +| # | Name space | Type | Defined at | +| ------ | ------------------------------------------------------------------ | ---------------------------------------------------- | ---------------------------------------------------------- | +| **N1** | `FieldOffset.value` — the *offset tag* | `str` (from `runtime.Offset.value: Union[int, str]`) | `iterator/runtime.py:36-37`, `ffront/fbuiltins.py:471-472` | +| **N2** | The **Python closure-variable name** the `FieldOffset` is bound to | `str` | consumed at `ffront/foast_to_gtir.py:305, 331` | +| **N3** | `Dimension.value` — the *dimension tag* | `str` | `common.py:79-81` | +| **N4** | `offset_provider` **dict key** | `str` | `common.py:1200-1209` | +| **N5** | ITIR `OffsetLiteral.value` / `AxisLiteral.value` | `str` | `iterator/ir.py:88-96` | + +`ts.OffsetType` carries **only** `source`/`target` +(`type_system/type_specifications.py:73-76`); N1 is discarded at +`fbuiltins.py:484-485`. That erasure is the root cause of most rows below. + +The only validation `FieldOffset` performs on itself is on the *kind*, never on +any name (`ffront/fbuiltins.py:470-482`): + +```python +class FieldOffset(runtime.Offset): # .value is the tag, inherited from runtime.Offset + source: common.Dimension + target: tuple[Dimension] | tuple[Dimension, Dimension] + + def __post_init__(self) -> None: + if len(self.target) == 2 and self.target[1].kind != common.DimensionKind.LOCAL: + raise ValueError("Second dimension in offset must be a local dimension.") +``` + +## 1. Concept inventory + +Every class and type alias in `src/gt4py/next/` that participates in the +offset/connectivity vocabulary, grouped by layer. + +### L0 — Vocabulary (`common.py`) + +| Concept | Where | Role | +| ------------------------------------------------- | ------------------------------ | -------------------------------------------------------------------------- | +| `Tag = str` | `common.py:63` | the alias that makes every name space stringly-typed | +| `DimensionKind` | `common.py:66-72` | `HORIZONTAL` / `VERTICAL` / `LOCAL` | +| `Dimension` | `common.py:79-81` | `(value: str, kind)`; `__add__`/`__sub__` build Cartesian shifts | +| `UnitRange`, `NamedRange`, `NamedIndex`, `Domain` | `common.py:196, 358, 369, 432` | index-space vocabulary; `Domain.dims` is where local dims appear on fields | + +### L1 — Connectivity objects (runtime data) + +| Concept | Where | Role | +| ----------------------------- | ----------------------------- | ------------------------------------------------------------------- | +| `Connectivity` | `common.py:990` | `Protocol`; a `Field` of indices with a `codomain` | +| `GatherConnectivity` | `common.py:1099` | nominal (not a `Protocol`): `premap` is a data-moving gather | +| `NeighborTable` | `common.py:1149` | `Protocol`; 2-D table-backed neighbor connectivity | +| `NdArrayConnectivityField` | `nd_array_field.py:516` | the concrete implementation | +| `NumPyArrayConnectivityField` | `nd_array_field.py:1032` | array-library variant | +| `CuPyArrayConnectivityField` | `nd_array_field.py:1049` | array-library variant | +| `JaxArrayConnectivityField` | `nd_array_field.py:1087` | array-library variant | +| `CartesianConnectivity` | `common.py:1241` | affine shift; no `ndarray`, **not** a `GatherConnectivity` | +| `StridedConnectivityField` | `iterator/embedded.py:107` | incomplete; iterator-view only (`TODO(havogt)`) | +| `_ConnectivityFileRef` | `otf/compilation_tasks.py:50` | lazy pickling stand-in; dumps to `.npy` to cross process boundaries | + +Constructors: `constructors.as_connectivity`, plus the `_field` / `_connectivity` +`singledispatch` pair at `common.py:1121-1146`. + +### L2 — Connectivity *types* (compile time) + +| Concept | Where | Contents | +| -------------------------- | ------------------- | --------------------------------------------------------------------------- | +| `ConnectivityType` | `common.py:964-973` | `domain`, `codomain`, `skip_value`, `dtype` | +| `NeighborConnectivityType` | `common.py:976-986` | `+ max_neighbors`; `.source_dim == domain[0]`, `.neighbor_dim == domain[1]` | + +### L3 — The provider (name to data binding) + +| Concept | Where | +| ------------------------------------------------------------------------ | --------------------- | +| `OffsetProvider = Mapping[Tag, NeighborTable]` | `common.py:1177` | +| `OffsetProviderType = Mapping[Tag, NeighborConnectivityType]` | `common.py:1178` | +| `OffsetProviderElem`, `OffsetProviderTypeElem` | `common.py:1172-1173` | +| `get_offset`, `get_offset_type`, `has_offset`, `offset_provider_to_type` | `common.py:1193-1221` | + +### L4 — Frontend declarations + +| Concept | Where | Contents | +| ---------------------------------- | --------------------------- | ---------------------------------------------- | +| `runtime.Offset` | `iterator/runtime.py:36-37` | `value: int \| str` | +| `FieldOffset(runtime.Offset)` | `ffront/fbuiltins.py:472` | `+ source`, `target` | +| `as_offset` builtin | `ffront/experimental.py:17` | dynamic Cartesian shift from an index field | +| `connectivity_for_cartesian_shift` | `common.py:1467` | builds a `CartesianConnectivity`; needs no tag | + +### L5 — Frontend types + +| Concept | Where | Note | +| ------------------- | ------------------------------ | ---------------------------------------------------- | +| `ts.OffsetType` | `type_specifications.py:73-79` | `source`/`target` only — **the tag is dropped here** | +| `ts.DimensionType` | `type_specifications.py:55` | wraps a `Dimension` | +| `ts.FieldType.dims` | `type_specifications.py:121` | sparse fields carry the local dim as a list member | + +### L6 — ITIR nodes + +| Node | Where | Carries | +| ---------------------- | ----------------------- | --------------------------------------------------- | +| `itir.OffsetLiteral` | `iterator/ir.py:88-89` | a bare `str` tag — unstructured | +| `itir.AxisLiteral` | `iterator/ir.py:92-96` | `value` + `kind` — a serialized `Dimension` | +| `itir.CartesianOffset` | `iterator/ir.py:99-101` | two `AxisLiteral`s — **no tag, no provider lookup** | + +### L7 — ITIR types + +| Concept | Where | Contents | +| --------------------------- | ------------------------------------------------ | ------------------------------------------------- | +| `it_ts.OffsetLiteralType` | `iterator/type_system/type_specifications.py:19` | `value: ScalarType \| str` | +| `it_ts.CartesianOffsetType` | `…:23` | `domain`, `codomain` | +| `it_ts.NamedRangeType` | `…:15` | `dim` | +| `it_ts.IteratorType` | `…:28` | `position_dims`, `defined_dims` | +| `ts.ListType` | `type_specifications.py:108-118` | `element_type` + `offset_type: Dimension \| None` | + +`ListType`'s docstring states the frontend/IR split explicitly: *"not used in the +frontend. The concept is represented as Field with local Dimension."* + +### L8 — Embedded iterator runtime + +| Concept | Where | Role | +| ---------------------------------- | --------------------------------- | --------------------------------------------------------- | +| `SparseTag(Tag)` | `iterator/embedded.py:102` | marks a shift into the sparse axis | +| `MDIterator`, `SparseListIterator` | `iterator/embedded.py:~800, 1507` | iterators; the latter holds `list_offset: Tag` | +| `_List`, `_ConstList` | `iterator/embedded.py:1399, 1420` | neighbor-list values | +| `_CONST_DIM` | `iterator/embedded.py:220` | reserved LOCAL dim, deliberately absent from the provider | +| position dicts | `iterator/embedded.py:597-616` | keyed by `Dimension.value` strings | + +### L9 — Backend representations + +| Concept | Where | Role | +| ------------------------------------------- | ------------------------------------------- | ----------------------------------------------------- | +| `gtfn_ir.OffsetLiteral` | `gtfn/gtfn_ir.py:52` | lowered tag | +| `gtfn_ir.TagDefinition` | `gtfn/gtfn_ir.py:249-251` | `name`, optional `alias`; emits `generated::_t` | +| `gtfn_ir.UnstructuredDomain.connectivities` | `gtfn/gtfn_ir.py:90-93` | `SymRef` to an offset declaration | +| `gtfn_ir.TaggedValues` | `gtfn/gtfn_ir.py:80-82` | tag-keyed sizes/offsets | +| `dace FieldopData` | `dace/lowering/gtir_to_sdfg_types.py:27-34` | carries the local-dim/offset-provider association | +| `dace connectivity_identifier` | `dace/sdfg_args.py:56-70` | `gt_conn_` array naming | + +### Concept count + +| Kind | Count | Notes | +| ------------------------------------------ | ----- | ---------------------------------------------------------------------------------------------------------------------- | +| Runtime connectivity classes | 8 | 3 are array-library variants of one; 1 is incomplete | +| Connectivity type classes | 2 | | +| Declaration classes | 2 | the subclassing is flagged as a conceptual mismatch at `fbuiltins.py:467` | +| Type-system representations of "an offset" | 5 | `ts.OffsetType`, `it_ts.OffsetLiteralType`, `it_ts.CartesianOffsetType`, `ts.ListType.offset_type`, `ts.DimensionType` | +| IR node kinds | 4 | 3 ITIR + 1 GTFN | +| Provider aliases | 4 | | + +Roughly **25 distinct concepts** for what is conceptually one thing — a mapping +between two index spaces — plus a name for it. + +## 2. How the concepts relate + +### 2.1 Connectivity class hierarchy + +```text +Field (Protocol) +└── Connectivity (Protocol) common.py:990 + ├── GatherConnectivity <- nominal, gather premap common.py:1099 + │ └── NeighborTable (Protocol, 2-D, table-backed) common.py:1149 + │ └── NdArrayConnectivityField nd_array_field.py:516 + │ ├── NumPyArrayConnectivityField nd_array_field.py:1032 + │ ├── CuPyArrayConnectivityField nd_array_field.py:1049 + │ └── JaxArrayConnectivityField nd_array_field.py:1087 + ├── CartesianConnectivity <- affine, no ndarray common.py:1241 + └── StridedConnectivityField <- WIP, iterator only iterator/embedded.py:107 +``` + +```mermaid +classDiagram + class Field { + <> + } + class Connectivity { + <> + +codomain: Dimension + +__gt_type__() ConnectivityType + } + class GatherConnectivity { + +ndarray + } + class NeighborTable { + <> + +__gt_type__() NeighborConnectivityType + } + class CartesianConnectivity { + +domain_dim + +offset: int + } + class StridedConnectivityField + class ConnectivityType { + +domain: tuple~Dimension~ + +codomain: Dimension + +skip_value + +dtype + } + class NeighborConnectivityType { + +max_neighbors: int + +source_dim + +neighbor_dim + } + Field <|-- Connectivity + Connectivity <|-- GatherConnectivity + Connectivity <|-- CartesianConnectivity + Connectivity <|-- StridedConnectivityField + GatherConnectivity <|-- NeighborTable + NeighborTable <|-- NdArrayConnectivityField + NdArrayConnectivityField <|-- NumPyArrayConnectivityField + NdArrayConnectivityField <|-- CuPyArrayConnectivityField + NdArrayConnectivityField <|-- JaxArrayConnectivityField + ConnectivityType <|-- NeighborConnectivityType + Connectivity ..> ConnectivityType : __gt_type__() + NeighborTable ..> NeighborConnectivityType : __gt_type__() +``` + +### 2.2 Declaration vs type vs data — the duplicated triple + +`FieldOffset` carries exactly the information in `NeighborConnectivityType` plus +a name, with **inverted vocabulary** and no cross-check. `FieldOffset.source` is +the connectivity's *codomain*; `FieldOffset.target` is its *domain*. The +inversion is because `source`/`target` describe the *field remap* (the field +lives on `source` and ends up on `target`), while `domain`/`codomain` describe +the *table*. + +```text + DECLARATION TYPE DATA + ─────────── ──── ──── + FieldOffset ts.OffsetType (none — bound later) + .value ─── dropped ──X + .target[0] ═══════════════ .target[0] ═══ A8 ═══════════ ConnectivityType.domain[0] + .target[1] ═══════════════ .target[1] ═══ A6 ═══════════ ConnectivityType.domain[1] + .source ═══════════════ .source ═══ A7 ═══════════ ConnectivityType.codomain + ^^^^^^^^^^^^^^^^^^^^^^^^^ + the same information, authored twice, + with inverted vocabulary, never cross-checked +``` + +### 2.3 Name flow — where the five name spaces diverge + +Four independently-authored strings converge on one dict lookup, and *which* of +them arrives there depends on the execution path and the operation. + +```text + ┌──────────────────────────────────────────────────┐ + │ V2EDim = Dimension("V2E", LOCAL) (N3) │ + USER AUTHORS │ V2E = FieldOffset("V2E", Edge,(V,V2EDim)) │ + FOUR STRINGS │ ^^^ (N2) │ + │ ^^^^^ (N1) │ + │ offset_provider = {"V2E": table} (N4) │ + └──────────────────────────────────────────────────┘ + │ + ┌─────────────────────────┴─────────────────────────┐ + │ │ + EMBEDDED PATH COMPILED PATH + │ │ + ┌──────────┴──────────┐ ┌─────────────┴─────────────┐ + │ shift │ reduce │ FOAST -> GTIR │ + │ fbuiltins.py:494 │ nd_array_field.py:983 │ foast_to_gtir.py:305,331 │ + │ uses N1 │ uses N3 (axis.value) │ uses N2 (Name.id) -> N5 │ + └──────────┬──────────┘ └─────────────┬─────────────┘ + │ │ + │ ┌──────────────────────────┤ + │ │ reduce: unroll_reduce.py:47 + │ │ uses N3 (ListType.offset_type.value) + │ │ + │ │ sparse arg: gtfn_module.py:95 + │ │ gtir_to_sdfg.py:581 + │ │ uses N3 (dim.value) + │ │ + └────────────┬───────────┴─────────────┬────────────┘ + v v + get_offset(offset_provider, ) == N4 + │ + v + NeighborTable / NeighborConnectivityType +``` + +```mermaid +flowchart TD + subgraph AUTHORED["User authors four strings"] + N3["N3 - Dimension('V2E', LOCAL)"] + N1["N1 - FieldOffset.value = 'V2E'"] + N2["N2 - python variable name V2E"] + N4["N4 - offset_provider key 'V2E'"] + end + + N1 --> EshiftE["embedded shift
fbuiltins.py:494"] + N3 --> EredE["embedded reduce
nd_array_field.py:983"] + N2 --> LOW["FOAST to GTIR
foast_to_gtir.py:305, 331"] + LOW --> N5["N5 - itir.OffsetLiteral"] + N5 --> CshiftC["compiled shift
type_synthesizer.py:748"] + N3 --> CredC["compiled reduce
unroll_reduce.py:47"] + N3 --> SPARSE["sparse field argument
gtfn_module.py:95
gtir_to_sdfg.py:581"] + + EshiftE --> GET + EredE --> GET + CshiftC --> GET + CredC --> GET + SPARSE --> GET + N4 -.->|"must equal the string that arrives"| GET + + GET["get_offset(offset_provider, string)
common.py:1200"] + GET --> DATA["NeighborTable / NeighborConnectivityType"] +``` + +### 2.4 The contrast that suggests the fix + +Cartesian shifts carry **dimensions** in the IR node; unstructured shifts carry +a **string** that must be resolved against a dict. Every constraint A1-A5 exists +only on the right-hand side. + +```text + CARTESIAN (already clean) UNSTRUCTURED (entangled) + ───────────────────────── ──────────────────────── + field(IDim + 1) field(V2E) + │ │ + v v + CartesianConnectivity itir.OffsetLiteral("V2E") <- a string + (common.py:1241) │ + │ v + v get_offset(provider, "V2E") + itir.CartesianOffset │ + domain: AxisLiteral v + codomain: AxisLiteral NeighborTable + (iterator/ir.py:99) + │ + v + NO tag. NO provider entry. NO lookup. +``` + +This is the concrete precedent behind any consolidation proposal: the Cartesian +path already eliminated the string indirection, and the unstructured path +retains it only because the neighbor table data must be supplied at runtime. + +## 3. Master table — cross-name-space identity constraints + +| # | Constraint | Embedded (field) | Embedded (iterator) | IR / type system | GTFN | DaCe | Enforced? | Source | +| ------- | ---------------------------------------------------------------------------- | ---------------- | ------------------- | ---------------- | ---------------- | ------------ | ----------------------------- | -------------------------------------------------------------------------------------------- | +| **A1** | `FieldOffset.value` (N1) == provider key (N4) | required | required | — | — | — | `KeyError` | `fbuiltins.py:494, 508`; `common.py:1207-1208` | +| **A2** | Python var name (N2) == provider key (N4) | — | — | required | required | required | silent; `KeyError` at runtime | `foast_to_gtir.py:305, 331` | +| **A3** | local dim `.value` (N3) == provider key (N4), **reductions** | required | required | required | required | required | `KeyError` | `nd_array_field.py:981-985`; `embedded.py:953, 1517, 1776`; `unroll_reduce.py:43-50, 61-65` | +| **A4** | local dim `.value` (N3) == provider key (N4), **sparse field args** | — | — | — | required | required | `assert` / `ValueError` | `gtfn_module.py:88-98`; `gtir_to_sdfg.py:572-585, 838-842`; `gtir_to_sdfg_lambda.py:766-770` | +| **A5** | `FieldOffset.value` (N1) == local dim `.value` (N3), **shift path** | n/a (A1 governs) | n/a | **not** required | **not** required | inconsistent | codegen branch handles it | `itir_to_gtfn_ir.py:181-190`; regression test | +| **A6** | `target[-1]` == connectivity `neighbor_dim` (full `Dimension` equality) | required | required | required | required | required | no eager check; index error | `fbuiltins.py:496`; `common.py:984-986` | +| **A7** | `FieldOffset.source` == connectivity `codomain` | required | required | required | required | required | `assert` only | `embedded.py:596-614`; `type_synthesizer.py:748-758` | +| **A8** | `FieldOffset.target[0]` == connectivity `domain[0]` (`source_dim`) | required | required | required | required | required | `assert found` | `type_synthesizer.py:752-758`; `embedded.py:597-599` | +| **A9** | `Dimension.value` (N3) is the key of the embedded **iterator position dict** | — | required | — | — | — | `assert ... in pos` | `embedded.py:574-576, 597-616, 941-950` | +| **A10** | `Dimension.value` (N3) round-trips through `AxisLiteral.value` (N5) | — | — | required | required | required | structural | `iterator/ir.py:92-96`; `ir_utils/misc.py:234-235`; `inference.py:463-464` | + +### Notes on A5 + +A5 is the only row with history. PR #1789 (`fix[next]: gtfn with offset name != local dimension name`) lifted it for shifts and added the +`if offset_name != connectivity_type.neighbor_dim.value` branch at +`itir_to_gtfn_ir.py:185-190`. Its regression test is +`tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py`, +whose docstring gives the motivation: + +> If the value of the `NeighborConnectivityType.neighbor_dim` did not match the +> `FieldOffset` value, gtfn would silently ignore the neighbor index, see +> . + +That test covers **only** `a(Off[1])` on `GTFN_CPU`. It does not cover +`neighbor_sum`, embedded execution, or DaCe. A3 and A4 were never lifted, so a +mismatch still breaks reductions and sparse arguments. + +The DaCe "inconsistent" entry: `gtir_to_sdfg_lambda.py:1155` builds +`Dimension(offset, LOCAL)` — a local dim named after the *tag* — while +`type_synthesizer.py:327-329` builds the same `ListType` from +`conn_type.neighbor_dim`. The two agree only when A5 holds. + +## 4. Per-context detail + +### 4.1 Embedded — field level (`nd_array_field`, `fbuiltins`) + +| Site | Key used | Constraint | +| ------------------------------------------------ | ----------------- | ------------------------------------------------------------------------------ | +| `fbuiltins.py:491-498` `FieldOffset.__getitem__` | `self.value` (N1) | A1; then `NamedIndex(self.target[-1], offset)` gives A6 | +| `fbuiltins.py:502-520` `as_connectivity_field` | `self.value` (N1) | A1 | +| `nd_array_field.py:981-985` reductions | `axis.value` (N3) | A3 — carries the comment `# assumes offset and local dimension have same name` | +| `nd_array_field.py:972-979` | — | `axis.kind == LOCAL`; at most one local dim per field | +| `nd_array_field.py:317-320` `premap` | — | `FieldOffset` to `Connectivity` via A1 | + +### 4.2 Embedded — iterator level (`iterator/embedded.py`) + +| Site | Key used | Constraint | +| --------------------------------------- | --------------------------------------------------------- | --------------------------------- | +| `:596-616` `execute_shift` | tag (N4), then `source_dim.value` / `codomain.value` (N3) | A7, A8, A9 | +| `:566-576` sparse shift | tag (N4) | A3 | +| `:941-953` `make_in_iterator` | `sparse_dimensions[0].value` (N3) used as tag | A3 | +| `:1517-1519` `SparseListIterator.deref` | `self.list_offset` (N3-derived) | A3 | +| `:1005` `field_setitem` | `value.offset.value` used as a **field dim name** | A3 (tag to N3, reverse direction) | +| `:1410-1416` `_List.__gt_type__` | tag, then `neighbor_dim` | correct direction, no assumption | +| `:1436-1451` `neighbors` | `offset.value` (N1) | A1 | +| `:1776` `_fieldspec_list_to_value` | `offset_type.value` (N3) | A3 | + +### 4.3 IR / type system + +| Site | Key used | Constraint | +| ------------------------------------------------------- | ----------------------------------------------- | ------------------------------------------------------------------------------------ | +| `type_synthesizer.py:326-329` `neighbors` | `OffsetLiteral.value` (N5), then `neighbor_dim` | A2; local dim taken from provider, **not** from the tag | +| `type_synthesizer.py:740-758` `shift` | N5, then `domain[0]`/`codomain` | A2, A7, A8 (`assert found`, `assert not found`) | +| `type_synthesizer.py:433-447` `_canonicalize_nb_fields` | field's LOCAL dim to `ListType.offset_type` | where N3 enters `ListType` and becomes an A3 key downstream | +| `type_synthesizer.py:546-556` `_resolve_dimensions` | N5, then `get_offset_type` | A2 | +| `unroll_reduce.py:43-50, 61-65` | `arg.type.offset_type.value` (N3) | **A3** | +| `domain_utils.py:205-223` | `off.value` (N5) | A2 | +| `pass_manager.py:55-63` | `source_dim.value`/`codomain.value` (N3) | domain sizes keyed by dimension name | +| `past_to_itir.py:409-410` | — | `ValueError: "common.Dimension '{dim.value}' must not be local."` in program domains | +| `type_deduction.py:459-464` | — | `"Second dimension in offset must be a local dimension."` | +| `type_info.py:637-650, 848-878` | — | shift typing via `source`/`target` only; the tag is never consulted | + +### 4.4 GTFN backend + +| Site | Name used | Constraint | +| ------------------------------------- | -------------------------------------------------- | ---------------------------------------------------------- | +| `itir_to_gtfn_ir.py:181-190` | provider key **and** `neighbor_dim.value` | the only site that anticipates A5 failing; emits both tags | +| `itir_to_gtfn_ir.py:191-196` | `source_dim.value`, `codomain.value` | must be `HORIZONTAL`, else `NotImplementedError` | +| `itir_to_gtfn_ir.py:197-200` | — | provider entries must be `NeighborConnectivityType` | +| `itir_to_gtfn_ir.py:485-492` | N5 tags | `o in self.offset_provider_type` | +| `itir_to_gtfn_ir.py:139-148, 166-180` | `dim.value` (N3) | every field dim name becomes a C++ tag | +| `gtfn_module.py:88-98` | `dim.value` (N3) | **A4** | +| `gtfn_module.py:126-136` | `domain[0].value`, `domain[1].value`, provider key | all three become `generated::_t` | + +### 4.5 DaCe backend + +| Site | Name used | Constraint | +| ---------------------------------------------------- | -------------------------------------- | ------------------------------------------------------------------------------------------------------------------- | +| `gtir_to_sdfg.py:572-585` | `local_dim.value` (N3) | **A4**, explicit: `ValueError("The provided local dimension {local_dim} does not match any offset provider type.")` | +| `gtir_to_sdfg.py:838-842` | `dim.value` (N3) | A4 — array shape from `max_neighbors` | +| `gtir_to_sdfg_lambda.py:766-770` | `local_dim.value` (N3) | A4 | +| `gtir_to_sdfg_lambda.py:1312-1319, 1371, 1443, 1455` | `offset_type.value` (N3) | A3, plus connectivity array name | +| `gtir_to_sdfg_lambda.py:1155` | tag (N5) to `Dimension(offset, LOCAL)` | reverse of A5; conflicts with `type_synthesizer.py:329` | +| `gtir_to_sdfg_lambda.py:1718-1727` | `offset_provider_arg.value` (N5) | genuine tag lookup — correct | +| `gtir_to_sdfg_primitives.py:324-331` | `offset_type.value` (N3) | A3 | +| `gtir_to_sdfg_scan.py:385-389` | `offset_type.value` (N3) | A3 | +| `sdfg_args.py:73-93` | field name plus `dim.value` | dim matched against `source_dim`/`neighbor_dim`, else `ValueError` | + +## 5. Constraints on the *format* of names + +| # | Constraint | Source | +| ------ | --------------------------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------------------------- | +| **F1** | `_Staggered` is a **reserved prefix**: any `Dimension` whose `value` starts with it is treated as staggered | `common.py:1444-1464` (`_STAGGERED_PREFIX = "_Staggered"`) | +| **F2** | GTFN aliases every staggered tag to its base tag by string surgery | `itir_to_gtfn_ir.py:703`, `_add_staggered_aliases:204-215` | +| **F3** | `_CONST_DIM` is a reserved LOCAL dimension name, deliberately *absent* from the provider and special-cased at every lookup | `embedded.py:220, 572, 1513, 1768`; `gtir_to_sdfg_lambda.py:62, 1314, 1355` | +| **F4** | **Dimension names determine memory layout** — `order_dimensions` sorts by `(kind, as_non_staggered(dim).value)` | `common.py:1334-1344` | +| **F5** | GTFN: every dim name and provider key becomes a C++ type `generated::_t`, so it must be a valid C++ identifier and collision-free | `gtfn_module.py:97, 130-136` | +| **F6** | GTFN connectivity params: `gt_conn_`, so keys must not collide **case-insensitively** | `gtfn_module.py:32, 120, 133` | +| **F7** | DaCe connectivity arrays: `gt_conn_`, recovered by regex `^gt_conn_(\S+)$` | `sdfg_args.py:24-25, 56-70` | +| **F8** | DaCe map variables: `i__gtx_[dim]`; map fusion/splitting transformations **rely on these strings matching** | `gtir_to_sdfg_utils.py:44-54` | +| **F9** | DaCe field symbols: `____size/stride`, `_range_symbol_name(field, dim.value)` | `sdfg_args.py:73-82, 119-122` | + +## 6. Structural (kind / arity) constraints + +| # | Constraint | Enforced | Source | +| --- | ------------------------------------------------------------------------- | ------------------------------------ | ----------------------------------------------------------------------------- | +| S1 | `len(target) == 2` implies `target[1].kind == LOCAL` | eager `ValueError` | `fbuiltins.py:480-482`; also `type_deduction.py:459-462` | +| S2 | A neighbor table's domain is exactly `(HORIZONTAL, LOCAL)` | `is_neighbor_table` guard | `common.py:1160-1168` | +| S3 | At most one LOCAL dim per field | `ValueError` / `NotImplementedError` | `common.py:1334-1337`; `nd_array_field.py:976-979`; `gtir_to_sdfg.py:586-589` | +| S4 | Cartesian offset iff `len(target)==1 and source==target[0]` and not LOCAL | predicate | `fbuiltins.py:524-529` | +| S5 | Non-Cartesian offset or LOCAL dim implies grid type `UNSTRUCTURED` | `ValueError` | `transform_utils.py:60-77` | +| S6 | `as_offset` is Cartesian-only | `DSLError` | `type_deduction.py:955-965` | +| S7 | Program domains must not contain LOCAL dims | `ValueError` | `past_to_itir.py:409-410` | + +## 7. Observed behaviour + +Two properties above were confirmed by running them, not only by reading. + +### 7.1 A1 vs A2 — embedded and compiled key on different strings + +```python +MyOff = gtx.FieldOffset("TAGNAME", source=E, target=(V, Neigh)) # tag != variable name + + +@gtx.field_operator +def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(MyOff[1]) +``` + +```text +embedded: offset_provider={"TAGNAME": conn} -> OK ; {"MyOff": conn} -> KeyError 'TAGNAME' +roundtrip: offset_provider={"MyOff": conn} -> OK ; {"TAGNAME": conn} -> KeyError 'MyOff' +``` + +The compiled path uses the Python variable name because `foast_to_gtir.py:302-306` +and `:325-331` emit `im.shift(offset_name.id, ...)` / `im.as_fieldop_neighbors(str(offset_name), ...)` +from the FOAST `Name.id` — never from `FieldOffset.value`. Lowering the operator +above yields: + +```text +foo = λ(a) → (⇑(λ(__it) → ·⟪MyOffₒ, 1ₒ⟫(__it)))(a); +``` + +### 7.2 A3 — reductions still require tag == local dim name + +Reusing the deliberately mismatched declaration from the #1789 regression test +(`Off` tagged `"Off"`, local dim named `"Neigh"`): + +```python +Off = gtx.FieldOffset("Off", source=E, target=(V, Neigh)) + + +@gtx.field_operator +def bar(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return neighbor_sum(a(Off), axis=Neigh) +``` + +```text +embedded: FAILED: KeyError: "Offset 'Neigh' not found in offset provider." +roundtrip: OK -> [30. 50. 40.] +``` + +`unroll_reduce.py:43-50` has the same assumption for the compiled pipeline +(established by reading; the roundtrip backend above does not exercise that pass). + +## 8. Practical consequence + +To be safe across **all** contexts, four strings must be identical: + +```text +FieldOffset.value == == offset_provider key == target[-1].value +``` + +plus `target[0] == conn.domain[0]` and `source == conn.codomain` as `Dimension` +objects (A6-A8). This is exactly what `tests/next_tests/toy_connectivity.py:18-26` +encodes: + +```python +V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) # value is "V2E", not "V2EDim" +V2E = gtx.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) +``` + +Relaxing any one of the four is currently supported only in the narrow slice +PR #1789 covered: shift-only, GTFN, no sparse arguments. Nothing validates the +full set up front — a violation surfaces as a `KeyError` from `common.py:1208`, +a bare `assert`, or, per the #1789 test docstring, silently wrong results. + +Two existing `TODO`s point at this tangle: + +- `common.py:976-977` — `NeighborConnectivityType`: *"refactor towards encoding + this information in the local dimensions of the `ConnectivityType.domain`"*. +- `fbuiltins.py:467-470` — *"`FieldOffset` and `runtime.Offset` are not an exact + conceptual match. Revisit if we want to continue subclassing here."* diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index eaf8aa8464..3ba57bae26 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -12,7 +12,9 @@ import collections import dataclasses import enum +import copyreg import functools +import importlib import math import re import sys @@ -124,55 +126,97 @@ def __str__(self) -> str: _DIM_KIND_ORDER = {DimensionKind.HORIZONTAL: 0, DimensionKind.LOCAL: 1, DimensionKind.VERTICAL: 2} -@dataclasses.dataclass(frozen=True) -class Dimension: - value: str - kind: DimensionKind = dataclasses.field(default=DimensionKind.HORIZONTAL) +class DimensionMeta(type): + """ + Metaclass of all dimension classes. - def __str__(self) -> str: - return f"{self.value}[{self.kind}]" + Holds the behaviour that used to live on `Dimension` *instances*, but on the class + object itself. Binary operators applied to a class object dispatch through its + metaclass, so this is the only place they can live. + """ + + kind: 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 + # `I == 5` overload. Without this every dimension class is unhashable, which breaks + # `domain({I: 2})`, dimension-keyed dicts, and eve's validator memoisation on annotation + # objects (so `ts.DimensionType` would fail at import). + __hash__ = type.__hash__ + + @property + def tag(cls) -> Tag: + """ + The dimension's identity: its qualified Python name, and its spelling in the IR. + + A property rather than a settable attribute, so it cannot drift from the type it + names. Use `__qualname__` for display; see `__str__`. + """ + return f"{cls.__module__}.{cls.__qualname__}" + + @property + def value(cls) -> NoReturn: + """ + Reject `SomeDim.value`, which used to be the dimension's name and is now `tag`. + + Without this the read silently returns the `value` slot descriptor of the *instance* + attribute rather than raising, and the nonsense value only surfaces much later -- as + a missing offset-provider key, or an `AxisLiteral` validation failure. Instance + access (`SomeDim(0).value`) is unaffected: a metaclass attribute is not on an + instance's lookup path. + """ + raise AttributeError( + f"'{cls.__qualname__}' is a dimension and has no 'value': its name is '.tag'," + f" and an *index* into it -- '{cls.__qualname__}(0)' -- is what has '.value'." + ) - def __call__(self, val: int) -> NamedIndex: - return NamedIndex(self, val) + def __repr__(cls) -> str: + return f"{cls.tag}[{cls.kind}]" - def __add__(self, offset: int | float) -> Connectivity: - return connectivity_for_cartesian_shift(self, offset) + 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}]" - def __sub__(self, offset: int | float) -> Connectivity: - return self + (-offset) + def __add__(cls: Dimension, offset: int | float) -> Connectivity: # type: ignore[misc] + return connectivity_for_cartesian_shift(cls, offset) - def __gt__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(value + 1, Infinity.POSITIVE),)) + def __sub__(cls: Dimension, offset: int | float) -> Connectivity: # type: ignore[misc] + return cls + (-offset) - def __ge__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(value, Infinity.POSITIVE),)) + def __gt__(cls: Dimension, value: core_defs.IntegralScalar) -> Domain: # type: ignore[misc] + return Domain(dims=(cls,), ranges=(UnitRange(value + 1, Infinity.POSITIVE),)) - def __lt__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(Infinity.NEGATIVE, value),)) + def __ge__(cls: Dimension, value: core_defs.IntegralScalar) -> Domain: # type: ignore[misc] + return Domain(dims=(cls,), ranges=(UnitRange(value, Infinity.POSITIVE),)) - def __le__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(Infinity.NEGATIVE, value + 1),)) + def __lt__(cls: Dimension, value: core_defs.IntegralScalar) -> Domain: # type: ignore[misc] + return Domain(dims=(cls,), ranges=(UnitRange(Infinity.NEGATIVE, value),)) - @overload # type: ignore[override] # incompatible with supertype `object.__eq__` which returns `bool`. - def __eq__(self, value: Dimension) -> bool: ... + def __le__(cls: Dimension, value: core_defs.IntegralScalar) -> Domain: # type: ignore[misc] + return Domain(dims=(cls,), ranges=(UnitRange(Infinity.NEGATIVE, value + 1),)) + + @overload # type: ignore[override] # incompatible with `type.__eq__`, which returns `bool`. + def __eq__(cls, value: DimensionMeta) -> bool: ... @overload - def __eq__(self, value: core_defs.IntegralScalar) -> Domain: ... - def __eq__(self, value: Dimension | core_defs.IntegralScalar) -> bool | Domain: - if isinstance(value, Dimension): - return self.value == value.value and self.kind == value.kind + def __eq__(cls, value: core_defs.IntegralScalar) -> Domain: ... + def __eq__( # type: ignore[misc] + cls: Dimension, value: DimensionMeta | core_defs.IntegralScalar + ) -> bool | Domain: + # NOTE: dimension-vs-dimension comparison is deliberately *not* handled here. A + # dimension's identity is its type, so `type.__eq__` (identity) is the correct + # answer; overriding it with `(tag, kind)` equality is what ADR 0028 rejects. if isinstance(value, core_defs.INTEGRAL_TYPES): - return Domain(dims=(self,), ranges=(UnitRange(value, value + 1),)) - # This will fallback to default identity comparison if reflection also returns `NotImplemented`, - # which does identity comparison, see https://docs.python.org/3/reference/datamodel.html#object.__eq__. + return Domain(dims=(cls,), ranges=(UnitRange(value, value + 1),)) return NotImplemented - @overload # type: ignore[override] # incompatible with supertype `object.__ne__` which returns `bool`. - def __ne__(self, value: Dimension) -> bool: ... + @overload # type: ignore[override] # incompatible with `type.__ne__`, which returns `bool`. + def __ne__(cls, value: DimensionMeta) -> bool: ... @overload - def __ne__(self, value: core_defs.IntegralScalar) -> Domain: ... - def __ne__(self, value: Dimension | core_defs.IntegralScalar) -> bool | Domain: - if isinstance(value, Dimension): - return self.value != value.value or self.kind != value.kind + def __ne__(cls, value: core_defs.IntegralScalar) -> Domain: ... + def __ne__( # type: ignore[misc] + cls: Dimension, value: DimensionMeta | core_defs.IntegralScalar + ) -> bool | Domain: if isinstance(value, core_defs.INTEGRAL_TYPES): raise NotImplementedError( "'Dimension.__ne__' with an integer value produces two disjoint domains, " @@ -182,26 +226,150 @@ def __ne__(self, value: Dimension | core_defs.IntegralScalar) -> bool | Domain: return NotImplemented -if TYPE_CHECKING: - # These exist as on-the fly replacements for Dimension instances - # (which are not types) during typechecking (with mypy). We can - # track up to four distinct dimensions at a time, everything beyond - # becomes AnyDim +class DimensionIndex(metaclass=DimensionMeta): + """ + A dimension. A concrete dimension is a *subclass*; an index along it an *instance*. + + This is the shape `enum.Enum` uses: the class is the collection, the instances are its + members. A dimension's identity is its type, and `tag` -- its qualified Python name -- + is how it is spelled in the IR. `value` is an index position along it. + + Examples: + >>> class I(DimensionIndex): ... + >>> class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... + >>> str(I), K.kind + ('I[horizontal]', ) + + >>> I(0) + I=0 + >>> I(0).dim is I, I(0).value + (True, 0) + + Two dimension classes are the same dimension only if they are the same class: + + >>> class I2(DimensionIndex): ... + >>> I == I2 + False + """ + + kind: ClassVar[DimensionKind] = DimensionKind.HORIZONTAL + + __slots__ = ("value",) + + #: Index position along the dimension. The dimension's *name* is `tag`, on the class. + value: int + + def __init_subclass__(cls, /, kind: Optional[DimensionKind] = None, **kwargs: Any) -> None: + super().__init_subclass__(**kwargs) + if "tag" in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}' sets 'tag' in its class body, which has no effect:" + " a dimension's tag is its qualified Python name. Rename the class instead." + ) + if "" in cls.__qualname__: + raise TypeError( + f"'{cls.__qualname__}' must be declared at module level: a dimension is" + " referenced from the IR by its qualified name, which has to be importable." + ) + if kind is not None: + cls.kind = kind + + def __init__(self, value: int) -> None: + self.value = value + + def __repr__(self) -> str: + return f"{type(self).__qualname__}={self.value}" + + __str__ = __repr__ + + def __eq__(self, other: object) -> bool: + if isinstance(other, DimensionIndex): + # NOTE: `is`, not `==`: a dimension's identity is its type (ADR 0028). + return type(self) is type(other) and self.value == other.value + return NotImplemented + + def __hash__(self) -> int: + return hash((type(self), self.value)) + + @property + def dim(self) -> Dimension: + """The dimension this index runs along, i.e. its own class.""" + return type(self) + + +#: A concrete dimension, i.e. the *class* itself rather than an index into it. +#: +#: NOTE: a PEP 695 `type` statement, not a plain `TypeAlias`, so that the removed +#: `Dimension("I")` spelling fails loudly. A plain alias for `type[X]` is a +#: `types.GenericAlias`, and calling one forwards to its `__origin__` while discarding the +#: arguments -- so `Dimension("I")` would evaluate to `type("I")`, i.e. `str`, with no error +#: at all. A `TypeAliasType` is simply not callable. +#: +#: The cost is that `get_origin()` of a PEP 695 alias is `None` rather than the aliased +#: origin, so a site dispatching on an annotation's shape must resolve it first (see +#: `xtyping.resolve_annotation`). `eve.datamodels` stores annotations *unresolved*, so this +#: applies to anything reading `__datamodel_fields__[...].type` too. See #2841 and ADR 0028. +type Dimension = type[DimensionIndex] + + +_STAGGERED_TAG_RE: Final = re.compile(r"^(?P[^\[\]]+)\[(?P.+)\]$") + + +@functools.cache +def resolve(tag: Tag) -> Dimension: + """ + Return the dimension class a tag names, by importing it. + + The counterpart of `DimensionMeta.tag`, for the IR boundaries that rebuild a dimension + from its name. A tag is a qualified Python name, so this is an import followed by an + attribute walk -- the same way `pickle` references a class. - @dataclasses.dataclass(frozen=True) - class _DimA(Dimension): ... + A purely dotted tag does not record where the module path ends and the qualname begins, + so the longest importable prefix wins and the remainder is walked as attributes. A + collision would need a module path and an attribute chain to have the same spelling. - @dataclasses.dataclass(frozen=True) - class _DimB(Dimension): ... + Parametrized dimensions such as `Staggered[K]` have no importable qualname; their tag + has the form `[]` and is resolved by subscripting the owner, which + goes through its intern table and so returns the identical class. - @dataclasses.dataclass(frozen=True) - class _DimC(Dimension): ... + Args: + tag: A dimension tag, as produced by `DimensionMeta.tag`. - @dataclasses.dataclass(frozen=True) - class _DimD(Dimension): ... + Returns: + The dimension class. - @dataclasses.dataclass(frozen=True) - class _AnyDim(Dimension): ... + Raises: + ValueError: If no prefix of `tag` is importable, or the attribute walk fails. + + Examples: + >>> resolve("gt4py.next.common.DimensionIndex") is DimensionIndex + True + """ + if (match := _STAGGERED_TAG_RE.match(tag)) is not None: + owner = resolve(match["owner"]) + return owner[resolve(match["base"])] # type: ignore[index] # parametrized dimension + + parts = tag.split(".") + for split in range(len(parts), 0, -1): + try: + obj: Any = importlib.import_module(".".join(parts[:split])) + except ImportError: + continue + for attr in parts[split:]: + try: + obj = getattr(obj, attr) + except AttributeError as ex: + raise ValueError( + f"Cannot resolve dimension tag '{tag}': '{'.'.join(parts[:split])}' has" + f" no attribute '{attr}'." + ) from ex + if not isinstance(obj, DimensionMeta): + raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") + return obj + raise ValueError( + f"Cannot resolve dimension tag '{tag}': no importable module prefix. A dimension" + " referenced from the IR must be declared at module level in an importable module." + ) class Infinity(enum.Enum): @@ -1490,20 +1658,153 @@ def __gt_builtin_func__(cls, /, func: fbuiltins.BuiltInFunction[_R, _P]) -> Call #: Equivalent to the `_FillValue` attribute in the UGRID Conventions #: (see: http://ugrid-conventions.github.io/ugrid-conventions/). _DEFAULT_SKIP_VALUE: Final[int] = -1 -_STAGGERED_PREFIX = "_Staggered" +#: Interned staggered dimensions, keyed by their *base dimension class*. +#: +#: NOTE: this is not the name-keyed dimension registry ADR 0028 rejects. It is memoization of +#: a type constructor -- keyed by identity, populated only by `StaggeredMeta.__getitem__`, and +#: never consulted to turn a user-authored name into a class. `typing`'s own subscription cache +#: plays the same role for generic aliases. +_STAGGERED_CACHE: dict[Dimension, Dimension] = {} + + +class StaggeredMeta(DimensionMeta): + """ + Metaclass of `Staggered`, whose subscription builds and interns a *real* class. + + A PEP 695 generic cannot be used here: `Staggered[K]` would be a `typing._GenericAlias`, + not a class, so it would fail `issubclass` and eve's `type[DimensionIndex]` validation, + and its `tag` could not name the base dimension. See ADR 0028. + """ + + #: Set by `__getitem__` on each parametrization. Its presence is what distinguishes a + #: staggered dimension from the bare `Staggered` base, which is also a `StaggeredMeta`. + base: Dimension + + def __str__(cls) -> str: + # NOTE: `__qualname__` has to carry the base's *full* tag so that `resolve` can find a + # base declared in another module, but that is too noisy for a diagnostic. Compose the + # short form here instead. + if "base" in cls.__dict__: + return f"Staggered[{cls.base.__qualname__}][{cls.kind}]" + return super().__str__() + + def __getitem__(cls, base: Dimension) -> Dimension: + if "base" in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}' is already staggered; a dimension cannot be staggered" + " twice." + ) + if not isinstance(base, DimensionMeta): + raise TypeError(f"'Staggered' expects a dimension, got '{base!r}'.") + if base not in _STAGGERED_CACHE: + _STAGGERED_CACHE[base] = cast( + Dimension, + StaggeredMeta( + f"Staggered[{base.__qualname__}]", + # NOTE: deliberately not `(cls, base)`. A staggered dimension is a + # *different* dimension, so `issubclass(Staggered[K], K)` must be false, + # or a staggered field would be accepted wherever a base one is required. + (cls,), + { + "_staggered_base": base, + "base": base, + "kind": base.kind, + "__slots__": (), + "__module__": cls.__module__, + "__qualname__": f"{cls.__qualname__}[{base.tag}]", + }, + ), + ) + return _STAGGERED_CACHE[base] + + +if TYPE_CHECKING: + # Checkers see an ordinary generic dimension, so `Staggered[K]` works in an annotation and + # inside `Field[Dims[Staggered[K]], ...]`. The runtime form below builds a real, interned + # class so that `issubclass` and eve's `type[...]` validation work. Verified clean under + # `mypy --strict` and pyright. + class Staggered[D: DimensionIndex](DimensionIndex): # noqa: D101 [undocumented-public-class] + base: ClassVar[Dimension] + +else: + + class Staggered(DimensionIndex, metaclass=StaggeredMeta): + """ + A dimension sitting at the half-integer positions of a base dimension (ADR 0026). + + `Staggered[K]` is a real, interned dimension class: subscripting the same base twice + returns the identical object, so it round-trips through the IR by identity. + """ + + __slots__ = () + + def __init_subclass__(cls, /, **kwargs: Any) -> None: + # NOTE: gate on the namespace marker the metaclass sets, not on a module-level + # "currently building" flag -- compilation runs in worker processes and threads. + if "_staggered_base" not in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}' cannot subclass a staggered dimension directly;" + " write 'Staggered[BaseDim]'." + ) + super().__init_subclass__(**kwargs) + + +class ConstListDim(DimensionIndex, kind=DimensionKind.LOCAL): + """ + The local dimension of a list whose length is known at compile time (`make_const_list`). + + Declared here, once, because it must be a *single* class. It used to be built + independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless + while dimensions compared by `(name, kind)` -- the two instances were equal. Under + nominal identity (ADR 0028) two declarations would be two different dimensions, and the + `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s + built by embedded execution. + + TODO: becomes an owner-less local dimension with an explicit size, generalising this from + length 1 to length *n*, once local dimensions know their connectivity. + """ + + __slots__ = () + + +def _reduce_staggered(cls: StaggeredMeta) -> Any: + """ + Pickle a staggered dimension through its base, falling back to by-reference. + + `Staggered[K].__qualname__` contains brackets, which `pickle.save_global` cannot look up, + and `copyreg` is the only hook consulted before `save_global` for a class. Reconstruction + goes through `StaggeredMeta.__getitem__`, so identity is preserved. + + The bare `Staggered` base is also a `StaggeredMeta` instance but has no `base`, so it must + fall through to ordinary by-reference pickling. + """ + if "base" not in cls.__dict__: + return cls.__qualname__ + return (_make_staggered, (cls.base,)) + + +def _make_staggered(base: Dimension) -> Dimension: + return Staggered[base] # type: ignore[valid-type] # runtime subscription, see StaggeredMeta + + +copyreg.pickle(StaggeredMeta, _reduce_staggered) # type: ignore[arg-type] # metaclass reducer def is_staggered(dim: Dimension) -> bool: - """Return whether `dim` is a staggered dimension.""" - return dim.value.startswith(_STAGGERED_PREFIX) + """ + Return whether `dim` is a staggered dimension. + + Checks for the marker the metaclass sets, not `issubclass(dim, Staggered)`: the latter is + also true of the bare `Staggered` base, which has no base dimension to recover. + """ + return "base" in dim.__dict__ def flip_staggered(dim: Dimension) -> Dimension: """Return the staggered counterpart of `dim`.""" if is_staggered(dim): - return Dimension(dim.value[len(_STAGGERED_PREFIX) :], dim.kind) - else: - return Dimension(f"{_STAGGERED_PREFIX}{dim.value}", dim.kind) + return cast(Dimension, dim.base) # type: ignore[attr-defined] # guarded by is_staggered + return Staggered[dim] # type: ignore[valid-type] # runtime subscription def as_non_staggered(dim: Dimension) -> Dimension: diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index c8ba839a8c..1b922556a8 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -214,7 +214,7 @@ def skip_value( # Magic local dimension for the result of a `make_const_list`. # A clean implementation will probably involve to tag the `make_const_list` # with the neighborhood it is meant to be used with. -_CONST_DIM = common.Dimension(value="_CONST_DIM", kind=common.DimensionKind.LOCAL) +_CONST_DIM = common.ConstListDim @runtime_checkable 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 1eb4900233..d05dcb8eaa 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 @@ -59,7 +59,10 @@ # Magic local dimension used for list of values with length known at compile-time. -_CONST_DIM: Final = gtx_common.Dimension(value="_CONST_DIM", kind=gtx_common.DimensionKind.LOCAL) +# NOTE: the canonical class from `common`, not a local declaration: under nominal identity a +# second declaration would be a *different* dimension and the `== _CONST_DIM` checks below +# would stop matching `ListType`s built by embedded execution. +_CONST_DIM: Final = gtx_common.ConstListDim @dataclasses.dataclass(frozen=True) From 183488fdaeea76fb5458e449af6ce52fc572db3e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 12:54:11 +0200 Subject: [PATCH 04/20] wip[next]: migrate the tree to dimension classes; unit tests green Continues the dimension-classes core. `tests/next_tests/unit_tests` is green (1752 passed); integration and DaCe suites not yet run. Source: * `.value` -> `.tag` on dimensions, replayed from #2844 by exact line match (88 sites) plus the lines that had drifted since its base. * `isinstance(x, Dimension)` and `case Dimension()` -> `DimensionMeta`, since `Dimension` is a PEP 695 alias that neither accepts. * `resolve()` at the IR boundaries that rebuild a dimension from a tag. * `NamedIndex` removed: an index is an instance of its dimension. Four sites tuple-unpacked an index and now read `.dim` / `.value`. * Injective mangling applied to generated identifiers in gtfn (the declared `TagDefinition`s and the references to them), nanobind, DaCe and the roundtrip emitter -- declarations and references must mangle identically, or the bindings name a C++ type that was never declared. * DaCe no longer synthesizes a local dimension from the offset tag; it uses the connectivity's own `neighbor_dim`. Display vs identity, now applied uniformly: diagnostics print `__qualname__`, IR and codegen use `tag`. `Staggered` overrides `tag` so its `__qualname__` can stay short. `order_dimensions` sorts by `__qualname__` too: sorting by the qualified tag made a field's canonical dimension order depend on *which module* declared each dimension (recorded in ADR 0028). Tests: * Every `Dimension(...)` declaration became a class (codemod); the 132 function-local ones were hoisted. No hoist changed behaviour: same-variable locals merge into one class, and every same-name/different-variable pair already differed in kind. * Inline `Dimension("X")` expressions map to one shared class per (name, kind) per file, because two such calls used to be equal. * One string per connectivity: a local dimension's tag is now qualified, so the `V2EDim = Dimension("V2E")` convention that made the offset tag, the local dimension and the provider key one string is gone. Restored symbolically -- `FieldOffset(V2EDim.tag, ...)`, `{V2EDim.tag: table}` -- so PR 4 changes only the declaration. * Cross-module duplicates: dimensions that were value-equal across modules are distinct now. `cases_utils` imports the six unstructured dimensions it shared with `toy_connectivity` instead of redeclaring them, and four tests import rather than redeclare. A same-*variable* check was not enough: the original *strings* decide equality, and `test_gtfn_module`'s `IDim` was `"I"` while `cases_utils`' was `"IDim"`, so those stay distinct. --- .../next/0028-Dimensions_As_Nominal_Types.md | 7 + src/gt4py/next/__init__.py | 6 + src/gt4py/next/common.py | 72 ++++--- src/gt4py/next/constructors.py | 20 +- src/gt4py/next/custom_layout_allocators.py | 8 +- src/gt4py/next/embedded/common.py | 30 ++- src/gt4py/next/embedded/nd_array_field.py | 17 +- src/gt4py/next/embedded/operators.py | 8 +- src/gt4py/next/ffront/fbuiltins.py | 2 +- .../ffront/foast_passes/type_deduction.py | 6 +- src/gt4py/next/ffront/foast_to_gtir.py | 6 +- src/gt4py/next/ffront/past_to_itir.py | 10 +- src/gt4py/next/ffront/transform_utils.py | 2 +- src/gt4py/next/ffront/type_info.py | 17 +- src/gt4py/next/iterator/embedded.py | 72 +++---- .../next/iterator/ir_utils/domain_utils.py | 6 +- src/gt4py/next/iterator/ir_utils/ir_makers.py | 6 +- src/gt4py/next/iterator/ir_utils/misc.py | 2 +- src/gt4py/next/iterator/tracing.py | 4 +- .../next/iterator/transforms/pass_manager.py | 4 +- .../iterator/transforms/remove_broadcast.py | 2 +- ...replace_get_domain_range_with_constants.py | 7 +- .../next/iterator/transforms/unroll_reduce.py | 2 +- .../next/iterator/type_system/inference.py | 2 +- .../iterator/type_system/type_synthesizer.py | 2 +- src/gt4py/next/otf/binding/nanobind.py | 2 +- src/gt4py/next/otf/compilation_tasks.py | 2 +- .../codegens/gtfn/codegen.py | 4 +- .../codegens/gtfn/gtfn_module.py | 6 +- .../codegens/gtfn/itir_to_gtfn_ir.py | 41 ++-- .../runners/dace/lowering/gtir_to_sdfg.py | 4 +- .../lowering/gtir_to_sdfg_concat_where.py | 2 +- .../dace/lowering/gtir_to_sdfg_lambda.py | 24 ++- .../dace/lowering/gtir_to_sdfg_primitives.py | 2 +- .../dace/lowering/gtir_to_sdfg_scan.py | 2 +- .../dace/lowering/gtir_to_sdfg_utils.py | 2 +- .../runners/dace/sdfg_args.py | 9 +- .../dace/transformations/loop_blocking.py | 2 +- .../dace/transformations/map_orderer.py | 2 +- .../program_processors/runners/roundtrip.py | 24 ++- src/gt4py/next/type_system/mypy_plugin.py | 94 +-------- src/gt4py/next/type_system/type_info.py | 6 +- .../next/type_system/type_specifications.py | 2 +- .../next/type_system/type_translation.py | 4 +- .../artifacts/custom_named_collections.py | 13 +- .../benchmarks/benchmark_program_call.py | 6 +- tests/next_tests/fixtures/past_common.py | 6 +- .../integration_tests/cases_utils.py | 44 ++-- .../dace_tests/test_orchestration.py | 6 +- ..._write_back_buffer_elimination_lowering.py | 21 +- .../ffront_tests/test_concat_where.py | 12 +- .../ffront_tests/test_external_local_field.py | 16 +- .../ffront_tests/test_foast_pretty_printer.py | 19 +- .../ffront_tests/test_named_collections.py | 4 +- .../ffront_tests/test_reductions.py | 56 ++--- .../ffront_tests/test_staggered.py | 3 +- .../test_temporaries_with_sizes.py | 11 +- .../feature_tests/ffront_tests/test_tuples.py | 4 +- .../ffront_tests/test_type_conversion.py | 2 +- .../instrumentation_tests/test_hooks.py | 2 +- .../iterator_tests/test_builtins.py | 11 +- .../iterator_tests/test_conditional.py | 2 +- .../iterator_tests/test_implicit_fencil.py | 3 +- .../iterator_tests/test_program.py | 4 +- .../test_strided_offset_provider.py | 16 +- .../iterator_tests/test_tuple.py | 11 +- .../ffront_tests/test_ffront_fvm_nabla.py | 6 +- .../ffront_tests/test_icon_like_scan.py | 4 - .../test_multiple_output_domains.py | 8 +- .../multi_feature_tests/fvm_nabla_setup.py | 19 +- .../iterator_tests/test_anton_toy.py | 10 +- .../iterator_tests/test_fvm_nabla.py | 33 +-- .../iterator_tests/test_if_stmt.py | 2 +- .../iterator_tests/test_temporaries.py | 7 +- .../test_with_toy_connectivity.py | 33 +-- .../embedded_tests/test_domain_pickle.py | 7 +- .../test_offset_dimensions_names.py | 28 ++- tests/next_tests/toy_connectivity.py | 37 ++-- tests/next_tests/unit_tests/conftest.py | 10 +- .../embedded_tests/test_basic_program.py | 2 +- .../unit_tests/embedded_tests/test_common.py | 30 +-- .../unit_tests/embedded_tests/test_context.py | 18 +- .../embedded_tests/test_nd_array_field.py | 196 ++++++++++-------- .../test_decorator_domain_deduction.py | 17 +- .../ffront_tests/test_diagnostic_messages.py | 4 +- .../unit_tests/ffront_tests/test_fbuiltins.py | 3 +- .../ffront_tests/test_foast_to_gtir.py | 43 ++-- .../ffront_tests/test_func_to_foast.py | 13 +- .../test_func_to_foast_error_line_number.py | 3 +- .../ffront_tests/test_past_to_gtir.py | 12 +- .../ffront_tests/test_source_utils.py | 16 +- .../unit_tests/ffront_tests/test_stages.py | 2 +- .../ffront_tests/test_type_deduction.py | 80 ++++--- .../ir_utils_tests/test_domain_utils.py | 74 ++++--- .../test_embedded_field_with_list.py | 22 +- .../iterator_tests/test_embedded_internals.py | 8 +- .../test_inline_dynamic_shifts.py | 5 +- .../iterator_tests/test_runtime_domain.py | 16 +- .../iterator_tests/test_type_inference.py | 45 ++-- .../transforms_tests/test_collapse_tuple.py | 4 +- ...t_concat_where_canonicalize_domain_args.py | 6 +- .../test_concat_where_expand_tuple_args.py | 6 +- ...st_concat_where_transform_to_as_fieldop.py | 8 +- .../transforms_tests/test_cse.py | 5 +- .../test_dead_code_elimination.py | 5 +- .../transforms_tests/test_domain_inference.py | 40 ++-- .../test_expand_tuple_maps.py | 4 +- .../transforms_tests/test_fuse_as_fieldop.py | 13 +- .../transforms_tests/test_global_tmps.py | 12 +- .../transforms_tests/test_inline_scalar.py | 5 +- .../transforms_tests/test_prune_casts.py | 4 +- .../test_prune_empty_concat_where.py | 20 +- .../transforms_tests/test_unroll_reduce.py | 71 ++++--- .../binding_tests/test_cpp_interface.py | 16 +- .../build_systems_tests/conftest.py | 29 ++- .../otf_tests/test_compiled_program.py | 2 +- .../unit_tests/otf_tests/test_runners.py | 23 +- .../gtfn_tests/test_gtfn_module.py | 6 +- .../runners_tests/dace_tests/test_dace.py | 4 +- .../dace_tests/test_dace_bindings.py | 4 +- .../dace_tests/test_gtir_to_sdfg.py | 79 +++---- .../transformation_tests/test_map_promoter.py | 13 +- tests/next_tests/unit_tests/test_common.py | 165 +++++++++------ .../unit_tests/test_constructors.py | 15 +- .../test_custom_layout_allocators.py | 72 +++++-- .../next_tests/unit_tests/test_field_utils.py | 5 +- tests/next_tests/unit_tests/test_utils.py | 10 +- .../type_system_tests/test_type_info.py | 51 +++-- .../test_type_translation.py | 9 +- 129 files changed, 1388 insertions(+), 957 deletions(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 08af4b4d24..93311b6b38 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -195,6 +195,13 @@ nested declaration it is already the readable form (`V2E.Local`). `repr()` shows the full tag and the kind, which disambiguates the rare case of two same-named dimensions from different modules appearing in one message. +Ordering follows the same rule, and it matters more than it looks: `order_dimensions` +determines a field's *canonical dimension order*. Keying it on `tag` would make +that order depend on **which module each dimension is declared in**, so moving a +declaration would silently reorder a field's dimensions. It is therefore keyed on +the unqualified name, with `tag` only breaking ties between same-named dimensions +from different modules so the order stays total. + `DimensionMeta` must declare `__hash__ = type.__hash__` explicitly: Python sets `__hash__ = None` on any class body defining `__eq__` without it, and `__eq__` stays for the `I == 5` → `Domain` overload. Without it every dimension class is diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index 3b9f97592f..e665024d7d 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -25,16 +25,19 @@ CartesianConnectivity, Connectivity, Dimension, + DimensionIndex, DimensionKind, Dims, Domain, Field, GridType, + Staggered, UnitRange, as_non_staggered, domain, flip_staggered, is_staggered, + resolve, unit_range, ) from .constructors import FieldConstructor, as_connectivity, as_field, empty, full, ones, zeros @@ -115,7 +118,10 @@ "is_scalar_type", # from common "Dimension", + "DimensionIndex", "DimensionKind", + "Staggered", + "resolve", "Dims", "Field", "CartesianConnectivity", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 3ba57bae26..b64bdd28a4 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -315,6 +315,24 @@ def dim(self) -> Dimension: _STAGGERED_TAG_RE: Final = re.compile(r"^(?P[^\[\]]+)\[(?P.+)\]$") +def staggered_base_tag(tag: Tag) -> Optional[Tag]: + """ + Return the base dimension's tag if `tag` names a staggered dimension, else `None`. + + Reads the `[]` grammar that `Staggered[D]` produces and `resolve` parses, so it + works on a bare tag without importing anything. That matters where tags of dimensions and + of offsets are mixed in one collection: an offset tag is not a dimension and cannot be + resolved, but it simply does not match. + + Examples: + >>> staggered_base_tag("gt4py.next.common.Staggered[pkg.KDim]") + 'pkg.KDim' + >>> staggered_base_tag("pkg.KDim") is None + True + """ + return match["base"] if (match := _STAGGERED_TAG_RE.match(tag)) is not None else None + + @functools.cache def resolve(tag: Tag) -> Dimension: """ @@ -583,20 +601,12 @@ def __str__(self) -> str: IntIndex: TypeAlias = int | core_defs.IntegralScalar -class NamedIndex(NamedTuple): - dim: Dimension - value: IntIndex - - def __str__(self) -> str: - return f"{self.dim}={self.value}" - - FiniteNamedRange: TypeAlias = NamedRange[FiniteUnitRange] RelativeIndexElement: TypeAlias = IntIndex | slice | types.EllipsisType -NamedSlice: TypeAlias = slice # once slice is generic we should do: slice[NamedIndex, NamedIndex, Literal[1]], see https://peps.python.org/pep-0696/ -AbsoluteIndexElement: TypeAlias = NamedIndex | NamedRange | NamedSlice +NamedSlice: TypeAlias = slice # once slice is generic we should do: slice[DimensionIndex, DimensionIndex, Literal[1]], see https://peps.python.org/pep-0696/ +AbsoluteIndexElement: TypeAlias = DimensionIndex | NamedRange | NamedSlice AnyIndexElement: TypeAlias = RelativeIndexElement | AbsoluteIndexElement -AbsoluteIndexSequence: TypeAlias = Sequence[NamedRange | NamedIndex] +AbsoluteIndexSequence: TypeAlias = Sequence[NamedRange | DimensionIndex] RelativeIndexSequence: TypeAlias = tuple[ slice | IntIndex | types.EllipsisType, ... ] # is a tuple but called Sequence for symmetry @@ -616,16 +626,16 @@ def is_finite_named_range(v: NamedRange) -> TypeGuard[FiniteNamedRange]: def is_named_slice(obj: AnyIndexSpec) -> TypeGuard[slice]: return isinstance(obj, slice) and ( - isinstance(obj.start, NamedIndex) and isinstance(obj.stop, NamedIndex) + isinstance(obj.start, DimensionIndex) and isinstance(obj.stop, DimensionIndex) ) def is_any_index_element(v: AnyIndexSpec) -> TypeGuard[AnyIndexElement]: - return is_int_index(v) or isinstance(v, (NamedRange, NamedIndex, slice)) or v is Ellipsis + return is_int_index(v) or isinstance(v, (NamedRange, DimensionIndex, slice)) or v is Ellipsis def is_absolute_index_sequence(v: AnyIndexSequence) -> TypeGuard[AbsoluteIndexSequence]: - return isinstance(v, Sequence) and all(isinstance(e, (NamedRange, NamedIndex)) for e in v) + return isinstance(v, Sequence) and all(isinstance(e, (NamedRange, DimensionIndex)) for e in v) def is_relative_index_sequence(v: AnyIndexSequence) -> TypeGuard[RelativeIndexSequence]: @@ -667,7 +677,7 @@ def __init__( ) assert dims is not None and ranges is not None # for mypy - if not all(isinstance(dim, Dimension) for dim in dims): + if not all(isinstance(dim, DimensionMeta) for dim in dims): raise ValueError( f"'dims' argument needs to be a 'tuple[Dimension, ...]', got '{dims}'." ) @@ -722,7 +732,7 @@ def __getitem__(self, index: slice) -> Self: ... def __getitem__(self, index: Dimension) -> NamedRange: ... def __getitem__(self, index: int | slice | Dimension) -> NamedRange | Domain: - if isinstance(index, Dimension): + if isinstance(index, DimensionMeta): try: index = self.dims.index(index) except ValueError as ex: @@ -851,7 +861,7 @@ def insert(self, index: int | Dimension, *named_ranges: NamedRange) -> Domain: def replace(self, index: int | Dimension, *named_ranges: NamedRange) -> Domain: assert all(isinstance(nr, NamedRange) for nr in named_ranges) - if isinstance(index, Dimension): + if isinstance(index, DimensionMeta): dim_index = self.dim_index(index) if dim_index is None: raise ValueError(f"Dimension '{index}' not found in Domain.") @@ -958,7 +968,7 @@ def __gt_domain__(self) -> Domain: @property def __gt_dims__(self) -> tuple[str, ...]: - return tuple(d.value for d in self.__gt_domain__.dims) + return tuple(d.tag for d in self.__gt_domain__.dims) @runtime_checkable @@ -1552,11 +1562,17 @@ 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'.") + # 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 + # the ordering a property of the dimensions themselves; `tag` only breaks ties between + # same-named dimensions from different modules, so the order stays total. return sorted( dims, key=lambda dim: ( _DIM_KIND_ORDER[dim.kind], - as_non_staggered(dim).value, + as_non_staggered(dim).__qualname__, + as_non_staggered(dim).tag, ), ) @@ -1586,7 +1602,7 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: 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 - lexicographically by `Dimension.value`. + lexicographically by `Dimension.tag`. Examples: >>> from gt4py.next.common import Dimension @@ -1680,13 +1696,15 @@ class StaggeredMeta(DimensionMeta): #: staggered dimension from the bare `Staggered` base, which is also a `StaggeredMeta`. base: Dimension - def __str__(cls) -> str: - # NOTE: `__qualname__` has to carry the base's *full* tag so that `resolve` can find a - # base declared in another module, but that is too noisy for a diagnostic. Compose the - # short form here instead. + @property + def tag(cls) -> Tag: + # NOTE: overridden so that `__qualname__` can stay the short, readable form used in + # diagnostics while the tag carries the base's *full* tag, which `resolve` needs to find + # a base declared in another module. Display is `__qualname__` and identity is `tag`, + # for staggered dimensions exactly as for any other. if "base" in cls.__dict__: - return f"Staggered[{cls.base.__qualname__}][{cls.kind}]" - return super().__str__() + return f"{cls.__module__}.Staggered[{cls.base.tag}]" + return super().tag def __getitem__(cls, base: Dimension) -> Dimension: if "base" in cls.__dict__: @@ -1711,7 +1729,7 @@ def __getitem__(cls, base: Dimension) -> Dimension: "kind": base.kind, "__slots__": (), "__module__": cls.__module__, - "__qualname__": f"{cls.__qualname__}[{base.tag}]", + "__qualname__": f"{cls.__qualname__}[{base.__qualname__}]", }, ), ) diff --git a/src/gt4py/next/constructors.py b/src/gt4py/next/constructors.py index e4320f99d3..cb20f79481 100644 --- a/src/gt4py/next/constructors.py +++ b/src/gt4py/next/constructors.py @@ -72,7 +72,7 @@ class FieldConstructor: def __init__( self, allocator: Allocator | None = None, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, device: core_defs.Device | None = None, ): if allocator is None: @@ -178,7 +178,7 @@ def as_field( ) -> nd_array_field.NdArrayField: """Create a `Field` from an array-like object. See :func:`as_field` for details.""" if isinstance(domain, Sequence) and all( - isinstance(dim, common.Dimension) for dim in domain + isinstance(dim, common.DimensionMeta) for dim in domain ): domain = cast(Sequence[common.Dimension], domain) if len(domain) != data.ndim: @@ -335,7 +335,7 @@ def asarray( class _CustomLayoutConstructor(_FieldArrayConstructor): allocator: next_allocators.FieldBufferAllocatorProtocol device: core_defs.Device | None = None - aligned_index: Sequence[common.NamedIndex] | None = None + aligned_index: Sequence[common.DimensionIndex] | None = None @functools.cached_property def device_id(self) -> int: @@ -384,7 +384,7 @@ def asarray( @eve.utils.optional_lru_cache def _field_constructor( allocator: Allocator | None, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, device: core_defs.Device | None = None, ) -> FieldConstructor: return FieldConstructor(allocator, aligned_index=aligned_index, device=device) @@ -395,7 +395,7 @@ def empty( domain: common.DomainLike, dtype: core_defs.DTypeLike = DEFAULT_DTYPE, *, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, allocator: Allocator | None = None, device: core_defs.Device | None = None, ) -> nd_array_field.NdArrayField: @@ -465,7 +465,7 @@ def zeros( domain: common.DomainLike, dtype: core_defs.DTypeLike = DEFAULT_DTYPE, *, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, allocator: Allocator | None = None, device: core_defs.Device | None = None, ) -> nd_array_field.NdArrayField: @@ -490,7 +490,7 @@ def ones( domain: common.DomainLike, dtype: core_defs.DTypeLike = DEFAULT_DTYPE, *, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, allocator: Allocator | None = None, device: core_defs.Device | None = None, ) -> nd_array_field.NdArrayField: @@ -516,7 +516,7 @@ def full( fill_value: core_defs.Scalar, dtype: core_defs.DTypeLike | None = None, *, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, allocator: Allocator | None = None, device: core_defs.Device | None = None, ) -> nd_array_field.NdArrayField: @@ -548,7 +548,7 @@ def as_field( dtype: core_defs.DTypeLike | None = None, *, origin: Mapping[common.Dimension, int] | None = None, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, allocator: Allocator | None = None, device: core_defs.Device | None = None, ) -> nd_array_field.NdArrayField: @@ -615,7 +615,7 @@ def as_connectivity( dtype: core_defs.DTypeLike | None = None, *, origin: Mapping[common.Dimension, int] | None = None, - aligned_index: Sequence[common.NamedIndex] | None = None, + aligned_index: Sequence[common.DimensionIndex] | None = None, allocator: Allocator | None = None, device: core_defs.Device | None = None, skip_value: core_defs.IntegralScalar | eve.NothingType | None = eve.NOTHING, diff --git a/src/gt4py/next/custom_layout_allocators.py b/src/gt4py/next/custom_layout_allocators.py index 8451863cc3..5eee3833ac 100644 --- a/src/gt4py/next/custom_layout_allocators.py +++ b/src/gt4py/next/custom_layout_allocators.py @@ -37,7 +37,7 @@ def __gt_allocate__( domain: common.Domain, dtype: core_defs.DType[core_defs.ScalarT], device_id: int = 0, - aligned_index: Optional[Sequence[common.NamedIndex]] = None, # absolute position + aligned_index: Optional[Sequence[common.DimensionIndex]] = None, # absolute position ) -> core_allocators.TensorBuffer[core_defs.DeviceTypeT, core_defs.ScalarT]: ... @@ -86,7 +86,7 @@ def is_field_allocation_tool_for( def _absolute_to_relative_index( - indices: Sequence[common.NamedIndex], domain: common.Domain + indices: Sequence[common.DimensionIndex], domain: common.Domain ) -> Sequence[int]: """Convert absolute indices to relative indices based on the domain's dimensions. @@ -130,7 +130,7 @@ def __gt_allocate__( domain: common.Domain, dtype: core_defs.DType[core_defs.ScalarT], device_id: int = 0, - aligned_index: Optional[Sequence[common.NamedIndex]] = None, # absolute position + aligned_index: Optional[Sequence[common.DimensionIndex]] = None, # absolute position ) -> core_allocators.TensorBuffer[core_defs.DeviceTypeT, core_defs.ScalarT]: shape = domain.shape layout_map = self.layout_mapper(domain.dims) @@ -214,7 +214,7 @@ def __gt_allocate__( domain: common.Domain, dtype: core_defs.DType[core_defs.ScalarT], device_id: int = 0, - aligned_index: Optional[Sequence[common.NamedIndex]] = None, # absolute position + aligned_index: Optional[Sequence[common.DimensionIndex]] = None, # absolute position ) -> core_allocators.TensorBuffer[core_defs.DeviceTypeT, core_defs.ScalarT]: raise self.exception diff --git a/src/gt4py/next/embedded/common.py b/src/gt4py/next/embedded/common.py index 829bd49c2d..07402b1079 100644 --- a/src/gt4py/next/embedded/common.py +++ b/src/gt4py/next/embedded/common.py @@ -70,7 +70,13 @@ def _absolute_sub_domain( for i, (dim, rng) in enumerate(domain): if (pos := _find_index_of_dim(dim, index)) is not None: named_idx = index[pos] - _, idx = named_idx + # NOTE: `NamedRange` is a namedtuple but an index is a `DimensionIndex` + # instance, so the payload has to be selected rather than unpacked. + idx = ( + named_idx.unit_range + if isinstance(named_idx, common.NamedRange) + else named_idx.value + ) if isinstance(idx, common.UnitRange): if not idx <= rng: raise embedded_exceptions.IndexOutOfBounds( @@ -144,9 +150,9 @@ def restrict_to_intersection( ) -def iterate_domain(domain: common.Domain) -> Iterator[tuple[common.NamedIndex]]: +def iterate_domain(domain: common.Domain) -> Iterator[tuple[common.DimensionIndex]]: for idx in itertools.product(*(list(r) for r in domain.ranges)): - yield tuple(common.NamedIndex(d, i) for d, i in zip(domain.dims, idx)) # type: ignore[misc] # trust me, `idx` is `tuple[int, ...]` + yield tuple(d(i) for d, i in zip(domain.dims, idx)) # type: ignore[misc] # trust me, `idx` is `tuple[int, ...]` def _expand_ellipsis( @@ -179,10 +185,12 @@ def _slice_range(input_range: common.UnitRange, slice_obj: slice) -> common.Unit def _find_index_of_dim( dim: common.Dimension, - domain_slice: common.Domain | Sequence[common.NamedRange | common.NamedIndex | Any], + domain_slice: common.Domain | Sequence[common.NamedRange | common.DimensionIndex | Any], ) -> Optional[int]: - for i, (d, _) in enumerate(domain_slice): - if dim == d: + # NOTE: `.dim`, not tuple unpacking: a `NamedRange` is a namedtuple, but an index is a + # `DimensionIndex` instance, and both expose `.dim`. + for i, elem in enumerate(domain_slice): + if dim == elem.dim: return i return None @@ -198,16 +206,16 @@ def canonicalize_any_index_sequence(index: common.AnyIndexSpec) -> common.AnyInd def _named_slice_to_named_range(idx: common.NamedSlice) -> common.NamedRange | common.NamedSlice: assert hasattr(idx, "start") and hasattr(idx, "stop") if common.is_named_slice(idx): - start_dim, start_value = idx.start - stop_dim, stop_value = idx.stop + start_dim, start_value = idx.start.dim, idx.start.value + stop_dim, stop_value = idx.stop.dim, idx.stop.value if start_dim != stop_dim: raise IndexError( - f"Dimensions slicing mismatch between '{start_dim.value}' and '{stop_dim.value}'." + f"Dimensions slicing mismatch between '{start_dim.__qualname__}' and '{stop_dim.__qualname__}'." ) assert isinstance(start_value, int) and isinstance(stop_value, int) return common.NamedRange(start_dim, common.UnitRange(start_value, stop_value)) - if isinstance(idx.start, common.NamedIndex) and idx.stop is None: + if isinstance(idx.start, common.DimensionIndex) and idx.stop is None: raise IndexError(f"Upper bound needs to be specified for {idx}.") - if isinstance(idx.stop, common.NamedIndex) and idx.start is None: + if isinstance(idx.stop, common.DimensionIndex) and idx.start is None: raise IndexError(f"Lower bound needs to be specified for {idx}.") return idx diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index fa08498671..916cccfb67 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -169,7 +169,7 @@ def from_array( assert issubclass(array.dtype.type, core_defs.SCALAR_TYPES) - assert all(isinstance(d, common.Dimension) for d in domain.dims), domain + assert all(isinstance(d, common.DimensionMeta) for d in domain.dims), domain assert len(domain) == array.ndim assert all(s == 1 or len(r) == s for r, s in zip(domain.ranges, array.shape)) @@ -533,11 +533,11 @@ def from_array( # type: ignore[override] assert issubclass(array.dtype.type, core_defs.INTEGRAL_TYPES) - assert all(isinstance(d, common.Dimension) for d in domain.dims), domain + assert all(isinstance(d, common.DimensionMeta) for d in domain.dims), domain assert len(domain) == array.ndim assert all(len(r) == s or s == 1 for r, s in zip(domain.ranges, array.shape)) - assert isinstance(codomain, common.Dimension) + assert isinstance(codomain, common.DimensionMeta) return cls(domain, array, codomain, _skip_value=skip_value) @@ -934,11 +934,11 @@ def _concat_where( def _as_offset(offset: fbuiltins.FieldOffset, offset_field: NdArrayField) -> common.Connectivity: if not fbuiltins.is_cartesian_offset(offset): - target_dims = ", ".join(d.value for d in offset.target) + target_dims = ", ".join(d.tag for d in offset.target) raise ValueError( f"'as_offset' is only supported for Cartesian offsets " f"(single target dimension equal to source dimension); " - f"got source '{offset.source.value}' and target ({target_dims})." + f"got source '{offset.source.__qualname__}' and target ({target_dims})." ) source_dim = offset.source coords = _identity_index_array( @@ -972,7 +972,7 @@ 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.value + current_offset_provider, axis.tag ) # assumes offset and local dimension have same name assert common.is_neighbor_table(offset_definition) new_domain = common.Domain(*[nr for nr in field.domain if nr.dim != axis]) @@ -1139,7 +1139,7 @@ def _astype(field: common.Field | core_defs.ScalarT | tuple, type_: type) -> NdA def _get_slices_from_domain_slice( domain: common.Domain, - domain_slice: common.Domain | Sequence[common.NamedRange | common.NamedIndex], + domain_slice: common.Domain | Sequence[common.NamedRange | common.DimensionIndex], ) -> common.RelativeIndexSequence: """Generate slices for sub-array extraction based on named ranges or named indices within a Domain. @@ -1159,7 +1159,8 @@ def _get_slices_from_domain_slice( for pos_old, (dim, _) in enumerate(domain): if (pos := embedded_common._find_index_of_dim(dim, domain_slice)) is not None: - _, index_or_range = domain_slice[pos] + elem = domain_slice[pos] + index_or_range = elem.unit_range if isinstance(elem, common.NamedRange) else elem.value slice_indices.append(_compute_slice(index_or_range, domain, pos_old)) else: slice_indices.append(slice(None)) diff --git a/src/gt4py/next/embedded/operators.py b/src/gt4py/next/embedded/operators.py index 3ad95c40df..2b0d1e7a24 100644 --- a/src/gt4py/next/embedded/operators.py +++ b/src/gt4py/next/embedded/operators.py @@ -64,10 +64,10 @@ def __call__( # type: ignore[override] assert isinstance(init_type, ts.TupleType | ts.ScalarType | ts.NamedCollectionType) res = field_utils.field_from_typespec(init_type, out_domain, xp) - def scan_loop(hpos: Sequence[common.NamedIndex]) -> None: + def scan_loop(hpos: Sequence[common.DimensionIndex]) -> None: acc: xtyping.MaybeNestedInTuple[core_defs.ScalarT] = self.init for k in scan_range.unit_range if self.forward else reversed(scan_range.unit_range): - pos = (*hpos, common.NamedIndex(scan_axis, k)) + pos = (*hpos, scan_axis(k)) new_args = [_tuple_at(pos, arg) for arg in args] new_kwargs = {k: _tuple_at(pos, v) for k, v in kwargs.items()} acc = self.fun(acc, *new_args, **new_kwargs) # type: ignore[arg-type] # need to express that the first argument is the same type as the return @@ -176,7 +176,7 @@ def _intersect_scan_args( def _tuple_assign_value( - pos: Sequence[common.NamedIndex], + pos: Sequence[common.DimensionIndex], target: xtyping.MaybeNestedInTuple[common.MutableField], source: xtyping.MaybeNestedInTuple[core_defs.Scalar], ) -> None: @@ -188,7 +188,7 @@ def impl(target: common.MutableField, source: core_defs.Scalar) -> None: def _tuple_at( - pos: Sequence[common.NamedIndex], + pos: Sequence[common.DimensionIndex], field: xtyping.MaybeNestedInTuple[common.Field | core_defs.Scalar], ) -> core_defs.Scalar | tuple[core_defs.ScalarT | tuple, ...]: @named_collections.tree_map_named_collection diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 2fd072a483..682af608ad 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -506,7 +506,7 @@ def __getitem__(self, offset: int) -> common.Connectivity: offset_definition = common.get_offset(current_offset_provider, self.value) assert common.is_neighbor_table(offset_definition) - named_index = common.NamedIndex(self.target[-1], offset) + named_index = (self.target[-1])(offset) connectivity = offset_definition[named_index] return connectivity diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 738b105a61..cc8fe1a52c 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -710,7 +710,7 @@ def _deduce_binop_type( raise errors.DSLError( right.location, f"Invalid offset '{right.value}' for a Cartesian shift of dimension " - f"'{left.type.dim.value}'.", + f"'{left.type.dim.__qualname__}'.", hints=[ ( "Use an integer offset to shift within the dimension, or a half-integer " @@ -989,12 +989,12 @@ def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: assert isinstance(arg_0, ts.OffsetType) assert isinstance(arg_1, ts.FieldType) if not fbuiltins.is_cartesian_offset(arg_0): - target_dims = ", ".join(d.value for d in arg_0.target) + target_dims = ", ".join(d.tag for d in arg_0.target) 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.value}' and target ({target_dims}).", + f"got source '{arg_0.source.__qualname__}' and target ({target_dims}).", ) if not type_info.is_integral(arg_1): raise errors.DSLError( diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 6913b5d6be..2b2553f814 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -232,12 +232,12 @@ def visit_Symbol(self, node: foast.Symbol, **kwargs: Any) -> itir.Sym: def visit_Name(self, node: foast.Name, **kwargs: Any) -> itir.SymRef | itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.value, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) return im.ref(node.id) def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.value, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) if isinstance(named_tup_type := node.value.type, ts.NamedCollectionType): ind = named_tup_type.keys.index(node.attr) @@ -309,7 +309,7 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: # `field(Dim + idx)` (where `idx` is integer or half integer) case foast.BinOp( op=dialect_ast_enums.BinaryOperator.ADD | dialect_ast_enums.BinaryOperator.SUB, - left=foast.LocatedNode(type=ts.DimensionType(dim=common.Dimension() as dim)), + left=foast.LocatedNode(type=ts.DimensionType(dim=common.DimensionMeta() as dim)), right=foast.Constant(value=offset_index), ): if arg.op == dialect_ast_enums.BinaryOperator.SUB: diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index cd908fca14..2ebdc8c74d 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.Dimension + all_closure_vars, fbuiltins.FieldOffset, common.DimensionMeta ) grid_type = transform_utils._deduce_grid_type( inp.data.grid_type, offsets_and_dimensions.values() @@ -174,7 +174,7 @@ def _column_axis(all_closure_vars: dict[str, Any]) -> Optional[common.Dimension] if len(scanops_per_axis.values()) != 1: scanops_per_axis_str = "\n".join( - f"- {dim.value}: {', '.join(scanops)}" for dim, scanops in scanops_per_axis.items() + f"- {dim.__qualname__}: {', '.join(scanops)}" for dim, scanops in scanops_per_axis.items() ) raise TypeError( @@ -382,7 +382,7 @@ def _construct_itir_domain_arg( for dim_i, dim in enumerate(out_type.dims): # an expression for the range of a dimension dim_range = im.call("get_domain_range")( - out_expr, itir.AxisLiteral(value=dim.value, kind=dim.kind) + out_expr, itir.AxisLiteral(value=dim.tag, kind=dim.kind) ) dim_start, dim_stop = im.tuple_get(0, dim_range), im.tuple_get(1, dim_range) @@ -407,11 +407,11 @@ def _construct_itir_domain_arg( ) if dim.kind == common.DimensionKind.LOCAL: - raise ValueError(f"common.Dimension '{dim.value}' must not be local.") + raise ValueError(f"common.Dimension '{dim.__qualname__}' must not be local.") domain_args.append( itir.FunCall( fun=itir.SymRef(id="named_range"), - args=[itir.AxisLiteral(value=dim.value, kind=dim.kind), lower, upper], + args=[itir.AxisLiteral(value=dim.tag, kind=dim.kind), lower, upper], ) ) diff --git a/src/gt4py/next/ffront/transform_utils.py b/src/gt4py/next/ffront/transform_utils.py index cac74ae88b..09c9d4b9ee 100644 --- a/src/gt4py/next/ffront/transform_utils.py +++ b/src/gt4py/next/ffront/transform_utils.py @@ -62,7 +62,7 @@ def _deduce_grid_type( if isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o): deduced_grid_type = common.GridType.UNSTRUCTURED break - if isinstance(o, common.Dimension) and o.kind == common.DimensionKind.LOCAL: + if isinstance(o, common.DimensionMeta) and o.kind == common.DimensionKind.LOCAL: deduced_grid_type = common.GridType.UNSTRUCTURED break diff --git a/src/gt4py/next/ffront/type_info.py b/src/gt4py/next/ffront/type_info.py index cfe1a51abb..900c0eeb05 100644 --- a/src/gt4py/next/ffront/type_info.py +++ b/src/gt4py/next/ffront/type_info.py @@ -8,7 +8,7 @@ import functools import inspect from collections.abc import Callable, Iterable -from typing import Any, Iterator, Sequence, cast +from typing import Final, Any, Iterator, Sequence, cast import gt4py.next.ffront.type_specifications as ts_ffront import gt4py.next.type_system.type_specifications as ts @@ -163,6 +163,19 @@ def _tree_map_type_constructor_drop_python_type( return result +class _UnknownDim(common.DimensionIndex): + """ + Placeholder for a dimension that cannot be determined, shown only in a diagnostic. + + Displays as `...` so the resulting error reads `Field[[...], ]`. It is never lowered + or resolved, so the unimportable qualname this gives its `tag` is harmless. + """ + + +_UnknownDim.__qualname__ = "..." +_UNKNOWN_DIM: Final = _UnknownDim + + def _scan_param_promotion( param: ts.TypeSpec, arg: ts.TypeSpec ) -> ts.FieldType | ts.TupleType | ts.NamedCollectionType: @@ -198,7 +211,7 @@ def _as_field(dtype: ts.TypeSpec, path: tuple[int, ...]) -> ts.FieldType: # argument type differ. As such we can not extract the dimensions # and just return a generic field shown in the error later on. # TODO: we want some generic field type here, but our type system does not support it yet. - return ts.FieldType(dims=[common.Dimension("...")], dtype=dtype) + return ts.FieldType(dims=[_UNKNOWN_DIM], dtype=dtype) # Note: In the promotion of the scalar type to field type we drop the information about # the original python type in NamedCollections as we want to be able to express compatibility diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 1b922556a8..6a21853f45 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -565,13 +565,13 @@ def execute_shift( for i, p in reversed(list(enumerate(new_entry))): # first shift applies to the last sparse dimensions of that axis type if p is None: - if tag == _CONST_DIM.value: + if tag == _CONST_DIM.tag: new_entry[i] = 0 else: offset_implementation = common.get_offset(offset_provider, tag) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().source_dim - cur_index = pos[source_dim.value] + cur_index = pos[source_dim.tag] assert common.is_int_index(cur_index) if offset_implementation[cur_index, index].as_scalar() in [ None, @@ -586,17 +586,17 @@ def execute_shift( if isinstance(tag, common.CartesianConnectivity): new_pos = copy.copy(pos) - value = new_pos.pop(tag.domain_dim.value) + value = new_pos.pop(tag.domain_dim.tag) assert common.is_int_index(value) - new_pos[tag.codomain.value] = value + index + tag.offset + new_pos[tag.codomain.tag] = value + index + tag.offset 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 - assert source_dim.value in pos + assert source_dim.tag in pos new_pos = pos.copy() - new_pos.pop(source_dim.value) - cur_index = pos[source_dim.value] + new_pos.pop(source_dim.tag) + cur_index = pos[source_dim.tag] assert common.is_int_index(cur_index) if offset_implementation[cur_index, index].as_scalar() in [ None, @@ -606,7 +606,7 @@ def execute_shift( else: new_index = offset_implementation[cur_index, index].as_scalar() assert new_index is not None - new_pos[offset_implementation.codomain.value] = int(new_index) + new_pos[offset_implementation.codomain.tag] = int(new_index) return new_pos @@ -895,7 +895,7 @@ def deref(self) -> Any: axes = _get_axes(self.field, ignore_zero_dims=True) if __debug__: - if not all(axis.value in shifted_pos.keys() for axis in axes if axis is not None): + if not all(axis.tag in shifted_pos.keys() for axis in axes if axis is not None): raise IndexError("Iterator position doesn't point to valid location for its field.") slice_column = dict[Tag, range]() if self.column_axis is not None: @@ -916,7 +916,7 @@ def _get_sparse_dimensions(axes: Sequence[common.Dimension]) -> list[common.Dime return [ axis for axis in axes - if isinstance(axis, common.Dimension) and axis.kind == common.DimensionKind.LOCAL + if isinstance(axis, common.DimensionMeta) and axis.kind == common.DimensionKind.LOCAL ] @@ -937,18 +937,18 @@ def make_in_iterator( new_pos: Position = pos.copy() for sparse_dim in set(sparse_dimensions): init = [None] * sparse_dimensions.count(sparse_dim) - new_pos[sparse_dim.value] = init # type: ignore[assignment] # looks like mypy is confused + new_pos[sparse_dim.tag] = init # type: ignore[assignment] # looks like mypy is confused if column_dimension is not None: column_range = embedded_context.get_closure_column_range().unit_range # if we deal with column stencil the column position is just an offset by which the whole column needs to be shifted assert column_range is not None - new_pos[column_dimension.value] = column_range.start + new_pos[column_dimension.tag] = column_range.start it = MDIterator( - inp, new_pos, column_axis=column_dimension.value if column_dimension is not None else None + inp, new_pos, column_axis=column_dimension.tag if column_dimension is not None else None ) if len(sparse_dimensions) >= 1: if len(sparse_dimensions) == 1: - return SparseListIterator(it, sparse_dimensions[0].value) + return SparseListIterator(it, sparse_dimensions[0].tag) else: raise NotImplementedError( f"More than one local dimension is currently not supported, got {sparse_dimensions}." @@ -974,9 +974,9 @@ def _translate_named_indices( self, _named_indices: NamedFieldIndices ) -> common.AbsoluteIndexSequence: named_indices: Mapping[common.Dimension, FieldIndex | SparsePositionEntry] = { - d: _named_indices[d.value] for d in self._ndarrayfield.__gt_domain__.dims + d: _named_indices[d.tag] for d in self._ndarrayfield.__gt_domain__.dims } - domain_slice: list[common.NamedRange | common.NamedIndex] = [] + domain_slice: list[common.NamedRange | common.DimensionIndex] = [] for d, v in named_indices.items(): if isinstance(v, range): domain_slice.append(common.NamedRange(d, common.UnitRange(v.start, v.stop))) @@ -985,10 +985,10 @@ def _translate_named_indices( assert common.is_int_index( v[0] ) # derefing a concrete element in a sparse field, not a slice - domain_slice.append(common.NamedIndex(d, v[0])) + domain_slice.append(d(v[0])) else: assert common.is_int_index(v) - domain_slice.append(common.NamedIndex(d, v)) + domain_slice.append(d(v)) return tuple(domain_slice) def field_getitem(self, named_indices: NamedFieldIndices) -> Any: @@ -1003,7 +1003,7 @@ def field_setitem(self, named_indices: NamedFieldIndices, value: Any): ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ - self._translate_named_indices({**named_indices, _CONST_DIM.value: 0}) + self._translate_named_indices({**named_indices, _CONST_DIM.tag: 0}) ] = value.value else: self._ndarrayfield[self._translate_named_indices(named_indices)] = value @@ -1037,13 +1037,13 @@ def get_ordered_indices(axes: Iterable[Axis], pos: NamedFieldIndices) -> tuple[F res.append(slice(None)) else: assert _is_field_axis(axis) - assert axis.value in pos - assert isinstance(axis.value, str) - elem = pos[axis.value] + assert axis.tag in pos + assert isinstance(axis.tag, str) + elem = pos[axis.tag] if _is_sparse_position_entry(elem): - sparse_position_tracker.setdefault(axis.value, 0) - res.append(elem[sparse_position_tracker[axis.value]]) - sparse_position_tracker[axis.value] += 1 + sparse_position_tracker.setdefault(axis.tag, 0) + res.append(elem[sparse_position_tracker[axis.tag]]) + sparse_position_tracker[axis.tag] += 1 else: assert isinstance(elem, (int, np.integer, slice, range)) res.append(elem) @@ -1154,9 +1154,9 @@ def premap( raise NotImplementedError() def restrict(self, item: common.AnyIndexSpec) -> Self: - if isinstance(item, Sequence) and all(isinstance(e, common.NamedIndex) for e in item): + if isinstance(item, Sequence) and all(isinstance(e, common.DimensionIndex) for e in item): assert len(item) == 1 - assert isinstance(item[0], common.NamedIndex) # for mypy errors on multiple lines below + assert isinstance(item[0], common.DimensionIndex) # for mypy errors on multiple lines below d, r = item[0] assert d == self._dimension assert isinstance(r, core_defs.INTEGRAL_TYPES) @@ -1507,7 +1507,7 @@ class SparseListIterator: offsets: Sequence[OffsetPart] = dataclasses.field(default_factory=list, kw_only=True) def deref(self) -> Any: - if self.list_offset == _CONST_DIM.value: + if self.list_offset == _CONST_DIM.tag: return _ConstList( value=self.it.shift(*self.offsets, SparseTag(self.list_offset), 0).deref() ) @@ -1651,7 +1651,7 @@ def impl(*iters: ItIterator): def _dimension_to_tag( domain: runtime.CartesianDomain | runtime.UnstructuredDomain, ) -> dict[Tag, range]: - return {k.value: v for k, v in domain.items()} + return {k.tag: v for k, v in domain.items()} def _validate_domain(domain: Domain, offset_provider_type: common.OffsetProviderType) -> None: @@ -1728,10 +1728,10 @@ def _extract_column_range(domain) -> common.NamedRange | eve.NothingType: col_range_placeholder.unit_range.is_empty() ) # check it's just the placeholder with empty range column_axis = col_range_placeholder.dim - if column_axis is not None and column_axis.value in domain: + if column_axis is not None and column_axis.tag in domain: return common.NamedRange( column_axis, - common.UnitRange(domain[column_axis.value].start, domain[column_axis.value].stop), + common.UnitRange(domain[column_axis.tag].start, domain[column_axis.tag].stop), ) return eve.NOTHING @@ -1747,7 +1747,7 @@ def _get_output_type( col_dim: Optional[common.Dimension] = None if isinstance(col_range, common.NamedRange): col_dim = col_range.dim - del domain[col_range.dim.value] + del domain[col_range.dim.tag] # determine dtype by computing result at one point pos_in_domain = next(iter(_domain_iterator(domain))) @@ -1769,8 +1769,8 @@ def _fieldspec_list_to_value( else: offset_provider = embedded_context.get_offset_provider() offset_type = type_.offset_type - assert isinstance(offset_type, common.Dimension) - connectivity = common.get_offset(offset_provider, offset_type.value) + assert isinstance(offset_type, common.DimensionMeta) + connectivity = common.get_offset(offset_provider, offset_type.tag) assert common.is_neighbor_table(connectivity) return domain.insert( len(domain), @@ -1824,7 +1824,7 @@ def closure( column_dim = None if isinstance(column_range, common.NamedRange): column_dim = column_range.dim - del domain[column_range.dim.value] + del domain[column_range.dim.tag] out = as_tuple_field(out) if is_tuple_of_field(out) else _wrap_field(out) promoted_ins = [promote_scalars(inp) for inp in ins] @@ -1840,7 +1840,7 @@ def closure( column_range = cast(common.NamedRange, column_range) col_pos = pos.copy() for k in column_range.unit_range: - col_pos[column_range.dim.value] = k + col_pos[column_range.dim.tag] = k assert _is_concrete_position(col_pos) out.field_setitem(col_pos, res[k]) # type: ignore[index] diff --git a/src/gt4py/next/iterator/ir_utils/domain_utils.py b/src/gt4py/next/iterator/ir_utils/domain_utils.py index 6bd28f6dbc..80330e8cc8 100644 --- a/src/gt4py/next/iterator/ir_utils/domain_utils.py +++ b/src/gt4py/next/iterator/ir_utils/domain_utils.py @@ -155,7 +155,7 @@ def from_expr(cls, node: itir.Node) -> SymbolicDomain: axis_literal, lower_bound, upper_bound = named_range.args assert isinstance(axis_literal, itir.AxisLiteral) - ranges[common.Dimension(value=axis_literal.value, kind=axis_literal.kind)] = ( + ranges[common.resolve(axis_literal.value)] = ( SymbolicRange(lower_bound, upper_bound) ) return cls(_GRID_TYPE_MAPPING[node.fun.id], ranges) @@ -222,10 +222,10 @@ def translate( new_dim = connectivity.codomain assert new_dim not in new_ranges or old_dim == new_dim - if symbolic_domain_sizes is not None and new_dim.value in symbolic_domain_sizes: + if symbolic_domain_sizes is not None and new_dim.tag in symbolic_domain_sizes: new_range = SymbolicRange( im.literal(str(0), builtins.INTEGER_INDEX_BUILTIN), - im.ensure_expr(symbolic_domain_sizes[new_dim.value]), + im.ensure_expr(symbolic_domain_sizes[new_dim.tag]), ) else: assert common.is_neighbor_table(connectivity) diff --git a/src/gt4py/next/iterator/ir_utils/ir_makers.py b/src/gt4py/next/iterator/ir_utils/ir_makers.py index cc7dc3ea3d..27dbaa5c50 100644 --- a/src/gt4py/next/iterator/ir_utils/ir_makers.py +++ b/src/gt4py/next/iterator/ir_utils/ir_makers.py @@ -90,7 +90,7 @@ def ensure_expr(expr_like: ExprLike) -> itir.Expr: return ref(expr_like) elif core_defs.is_scalar_type(expr_like): return literal_from_value(expr_like) - elif isinstance(expr_like, common.Dimension): + elif isinstance(expr_like, common.DimensionMeta): return axis_literal(expr_like) assert isinstance(expr_like, itir.Expr), expr_like return expr_like @@ -583,7 +583,7 @@ def _impl(*its: itir.Expr) -> itir.FunCall: def axis_literal(dim: common.Dimension) -> itir.AxisLiteral: - return itir.AxisLiteral(value=dim.value, kind=dim.kind) + return itir.AxisLiteral(value=dim.tag, kind=dim.kind) def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: @@ -640,7 +640,7 @@ def index(dim: common.Dimension) -> itir.FunCall: Returns: A function that constructs a Field of indices in the given dimension. """ - return call("index")(itir.AxisLiteral(value=dim.value, kind=dim.kind)) + return call("index")(itir.AxisLiteral(value=dim.tag, kind=dim.kind)) def map_list(op): diff --git a/src/gt4py/next/iterator/ir_utils/misc.py b/src/gt4py/next/iterator/ir_utils/misc.py index d9a6e85f57..65fc047f5f 100644 --- a/src/gt4py/next/iterator/ir_utils/misc.py +++ b/src/gt4py/next/iterator/ir_utils/misc.py @@ -232,7 +232,7 @@ def grid_type_from_domain(domain: itir.FunCall) -> common.GridType: def dim_from_axis_literal(axis_literal: itir.AxisLiteral) -> common.Dimension: - return common.Dimension(value=axis_literal.value, kind=axis_literal.kind) + return common.resolve(axis_literal.value) def _flatten_tuple_expr(expr: itir.Expr) -> tuple[itir.Expr]: diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 4450d1c276..4e5b8d4b5a 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -140,8 +140,8 @@ def __call__(self, *args): def make_node(o): if isinstance(o, Node): return o - if isinstance(o, common.Dimension): - return AxisLiteral(value=o.value, kind=o.kind) + if isinstance(o, common.DimensionMeta): + return AxisLiteral(value=o.tag, kind=o.kind) if isinstance(o, common.Infinity): if o is common.Infinity.POSITIVE: return itir.InfinityLiteral.POSITIVE diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index 067dde468f..ec4363aba3 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -55,8 +55,8 @@ 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.value - codomain_dim = provider.__gt_type__().codomain.value + src_dim = provider.__gt_type__().source_dim.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( sizes.get(codomain_dim, 0), diff --git a/src/gt4py/next/iterator/transforms/remove_broadcast.py b/src/gt4py/next/iterator/transforms/remove_broadcast.py index a0c46b21de..52c61105a4 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -30,7 +30,7 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): >>> expr = im.call("broadcast")( ... im.ref("inp"), ... im.make_tuple( - ... *(itir.AxisLiteral(value=dim.value, kind=dim.kind) for dim in (IDim, JDim)) + ... *(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (IDim, JDim)) ... ), ... ) >>> expr.annex.domain = domain_utils.SymbolicDomain.from_expr(domain) diff --git a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py index f3696a9f80..868f9c81b6 100644 --- a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py +++ b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py @@ -114,7 +114,10 @@ def visit_FunCall(self, node: itir.FunCall, **kwargs) -> itir.FunCall: f"'{field}'." ) - index = next((i for i, d in enumerate(domain.dims) if d.value == dim.value), None) - assert index is not None, f"Dimension {dim.value} not found in {domain.dims}" + # NOTE: `dim` is the IR argument -- an `AxisLiteral` whose `value` is the dimension's tag + # -- while `domain.dims` holds dimension classes, so the tag is what they share. + assert isinstance(dim, itir.AxisLiteral) + index = next((i for i, d in enumerate(domain.dims) if d.tag == dim.value), None) + assert index is not None, f"Dimension '{dim.value}' not found in {domain.dims}" return im.make_tuple(domain.ranges[index].start, domain.ranges[index].stop) diff --git a/src/gt4py/next/iterator/transforms/unroll_reduce.py b/src/gt4py/next/iterator/transforms/unroll_reduce.py index c4ea32fc36..cfb7bb3226 100644 --- a/src/gt4py/next/iterator/transforms/unroll_reduce.py +++ b/src/gt4py/next/iterator/transforms/unroll_reduce.py @@ -44,7 +44,7 @@ def _get_partial_offset_tags(reduce_args: Iterable[itir.Expr]) -> Iterable[str]: assert all(isinstance(arg.type, ts.ListType) for arg in reduce_args) return [ - arg.type.offset_type.value # type: ignore[union-attr] # checked in previous lines + arg.type.offset_type.tag # 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 ] diff --git a/src/gt4py/next/iterator/type_system/inference.py b/src/gt4py/next/iterator/type_system/inference.py index c2ee562fff..b878640f12 100644 --- a/src/gt4py/next/iterator/type_system/inference.py +++ b/src/gt4py/next/iterator/type_system/inference.py @@ -462,7 +462,7 @@ def visit_SetAt(self, node: itir.SetAt, *, ctx) -> None: assert target_type.dtype == expr_type.dtype def visit_AxisLiteral(self, node: itir.AxisLiteral, **kwargs) -> ts.DimensionType: - return ts.DimensionType(dim=common.Dimension(value=node.value, kind=node.kind)) + return ts.DimensionType(dim=common.resolve(node.value)) # TODO: revisit what we want to do with OffsetLiterals as we already have an Offset type in # the frontend. diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index abaaaf27a2..a1622b0351 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -594,7 +594,7 @@ def applied_as_fieldop( ), ) - assert all(isinstance(dim, common.Dimension) for dim in output_dims) + assert all(isinstance(dim, common.DimensionMeta) for dim in output_dims) deduced_domain = ts.DomainType(dims=output_dims) if deduced_domain: diff --git a/src/gt4py/next/otf/binding/nanobind.py b/src/gt4py/next/otf/binding/nanobind.py index ab4004d44e..1afddc992f 100644 --- a/src/gt4py/next/otf/binding/nanobind.py +++ b/src/gt4py/next/otf/binding/nanobind.py @@ -211,7 +211,7 @@ def make_argument( source_buffer=name, dimensions=[ DimensionSpec( - name=dim.value, + name=common.codegen_name(dim.tag), static_stride=1 if ( unstructured_horizontal_has_unit_stride diff --git a/src/gt4py/next/otf/compilation_tasks.py b/src/gt4py/next/otf/compilation_tasks.py index ace3127501..6427c8ecc0 100644 --- a/src/gt4py/next/otf/compilation_tasks.py +++ b/src/gt4py/next/otf/compilation_tasks.py @@ -117,7 +117,7 @@ def _offset_provider_with_file_refs( ) -> common.OffsetProvider: return { name: value - if isinstance(value, common.Dimension) + if isinstance(value, common.DimensionMeta) else typing.cast(common.OffsetProviderElem, _ConnectivityFileRef(value)) for name, value in offset_provider.items() } diff --git a/src/gt4py/next/program_processors/codegens/gtfn/codegen.py b/src/gt4py/next/program_processors/codegens/gtfn/codegen.py index 3c7750184c..f44cffda4e 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/codegen.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/codegen.py @@ -145,7 +145,9 @@ def visit_TaggedValues(self, node: gtfn_ir.TaggedValues, **kwargs: Any) -> str: ) def visit_OffsetLiteral(self, node: gtfn_ir.OffsetLiteral, **kwargs: Any) -> str: - return node.value if isinstance(node.value, str) else f"{node.value}_c" + # NOTE: a string offset literal names a tag type declared as `generated::_t`, so + # it must be mangled exactly as the declaration was (see `common.codegen_name`). + return common.codegen_name(node.value) if isinstance(node.value, str) else f"{node.value}_c" SidComposite = as_mako( "::gridtools::sid::composite::keys<${','.join(f'::gridtools::integral_constant' for i in range(len(values)))}>::make_values(${','.join(values)})" 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 f861d4b182..5bc10c2fb5 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -89,7 +89,7 @@ def _process_regular_arguments( or dim.kind is common.DimensionKind.LOCAL ): # translate sparse dimensions to tuple dtype - dim_name = dim.value + dim_name = dim.tag connectivity = common.get_offset_type(offset_provider_type, dim_name) assert isinstance(connectivity, common.NeighborConnectivityType) size = connectivity.max_neighbors @@ -124,8 +124,8 @@ def _process_connectivity_args( # connectivity argument expression nbtbl = ( f"gridtools::fn::sid_neighbor_table::as_neighbor_table<" - f"generated::{connectivity_type.domain[0].value}_t, " - f"generated::{connectivity_type.domain[1].value}_t, " + f"generated::{common.codegen_name(connectivity_type.domain[0].tag)}_t, " + f"generated::{common.codegen_name(connectivity_type.domain[1].tag)}_t, " f"{connectivity_type.max_neighbors}" f">(std::forward({GENERATED_CONNECTIVITY_PARAM_PREFIX}{name.lower()}))" ) 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 31b64c78e2..0aec096134 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 @@ -107,7 +107,7 @@ def _collect_dimensions_from_domain( for nr in domain.args: assert isinstance(nr, itir.FunCall) dim_name = _name_from_named_range(nr) - offset_definitions[dim_name] = TagDefinition(name=Sym(id=dim_name)) + offset_definitions[dim_name] = TagDefinition(name=Sym(id=common.codegen_name(dim_name))) elif domain.fun == itir.SymRef(id="unstructured_domain"): if len(domain.args) > 2: raise ValueError("Unstructured_domain must not have more than 2 arguments.") @@ -116,14 +116,14 @@ def _collect_dimensions_from_domain( assert isinstance(horizontal_range, itir.FunCall) horizontal_name = _name_from_named_range(horizontal_range) offset_definitions[horizontal_name] = TagDefinition( - name=Sym(id=horizontal_name), alias=_horizontal_dimension + name=Sym(id=common.codegen_name(horizontal_name)), alias=_horizontal_dimension ) if len(domain.args) > 1: vertical_range = domain.args[1] assert isinstance(vertical_range, itir.FunCall) vertical_name = _name_from_named_range(vertical_range) offset_definitions[vertical_name] = TagDefinition( - name=Sym(id=vertical_name), alias=_vertical_dimension + name=Sym(id=common.codegen_name(vertical_name)), alias=_vertical_dimension ) else: raise AssertionError( @@ -148,7 +148,7 @@ def _collect_dimensions_from_params( for type_ in type_info.primitive_constituents(param.type): if isinstance(type_, ts.FieldType): for dim in type_.dims: - offset_definitions[dim.value] = TagDefinition(name=Sym(id=dim.value)) + offset_definitions[dim.tag] = TagDefinition(name=Sym(id=common.codegen_name(dim.tag))) return offset_definitions @@ -170,31 +170,31 @@ def _collect_offset_definitions( ] for dim in dims: if grid_type == common.GridType.CARTESIAN: - offset_definitions[dim.value] = TagDefinition(name=Sym(id=dim.value)) + offset_definitions[dim.tag] = TagDefinition(name=Sym(id=common.codegen_name(dim.tag))) else: assert grid_type == common.GridType.UNSTRUCTURED if dim.kind != common.DimensionKind.VERTICAL: raise ValueError( "Mapping an offset to a horizontal dimension in unstructured is not allowed." ) - offset_definitions[dim.value] = TagDefinition( - name=Sym(id=dim.value), alias=_vertical_dimension + offset_definitions[dim.tag] = TagDefinition( + name=Sym(id=common.codegen_name(dim.tag)), alias=_vertical_dimension ) for offset_name, connectivity_type in offset_provider_type.items(): if isinstance(connectivity_type, common.NeighborConnectivityType): assert grid_type == common.GridType.UNSTRUCTURED - offset_definitions[offset_name] = TagDefinition(name=Sym(id=offset_name)) - if offset_name != connectivity_type.neighbor_dim.value: - offset_definitions[connectivity_type.neighbor_dim.value] = TagDefinition( - name=Sym(id=connectivity_type.neighbor_dim.value) + 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)) ) for dim in [connectivity_type.source_dim, connectivity_type.codomain]: if dim.kind != common.DimensionKind.HORIZONTAL: raise NotImplementedError() - offset_definitions[dim.value] = TagDefinition( - name=Sym(id=dim.value), alias=_horizontal_dimension + offset_definitions[dim.tag] = TagDefinition( + name=Sym(id=common.codegen_name(dim.tag)), alias=_horizontal_dimension ) else: raise AssertionError( @@ -210,11 +210,12 @@ def _add_staggered_aliases( result: dict[str, TagDefinition] = {} aliases: dict[str, TagDefinition] = {} for name, tag_def in offset_definitions.items(): - if tag_def.alias is None and common.is_staggered(common.Dimension(value=name)): - base_name = common.as_non_staggered(common.Dimension(value=name)).value + # NOTE: read from the tag grammar, not by resolving `name`: this dict mixes dimension + # tags with offset tags, and an offset tag is not a dimension and cannot be resolved. + if tag_def.alias is None and (base_name := common.staggered_base_tag(name)) is not None: # ensure the base tag exists (as alias target and loop dimension) in this position - result.setdefault(base_name, TagDefinition(name=Sym(id=base_name))) - aliases[name] = TagDefinition(name=Sym(id=name), alias=SymRef(id=base_name)) + result.setdefault(base_name, TagDefinition(name=Sym(id=common.codegen_name(base_name)))) + aliases[name] = TagDefinition(name=Sym(id=common.codegen_name(name)), alias=SymRef(id=common.codegen_name(base_name))) else: result[name] = tag_def return {**result, **aliases} @@ -406,7 +407,7 @@ def visit_CartesianOffset(self, node: itir.CartesianOffset, **kwargs: Any) -> Li def visit_AxisLiteral(self, node: itir.AxisLiteral, **kwargs: Any) -> Literal: assert isinstance(node.type, ts.DimensionType) - return Literal(value=node.type.dim.value, type="axis_literal") + return Literal(value=node.type.dim.tag, type="axis_literal") def _make_domain(self, node: itir.FunCall) -> tuple[TaggedValues, TaggedValues]: tags = [] @@ -673,12 +674,12 @@ def convert_el_to_sid(el_expr: Expr, el_type: ts.ScalarType | ts.FieldType) -> E init=self.visit(stencil.args[2], **kwargs), ) column_axis = self.column_axis - assert isinstance(column_axis, common.Dimension) + assert isinstance(column_axis, common.DimensionMeta) return ScanExecution( backend=backend, scans=[scan], args=[self._visit_output_argument(node.target), *lowered_inputs], - axis=SymRef(id=column_axis.value), + axis=SymRef(id=column_axis.tag), ) assert projector is None # only scans have projectors return StencilExecution( 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 e645068f64..c4b1526bdc 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 @@ -578,7 +578,7 @@ 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.value): + if not gtx_common.has_offset(self.offset_provider_type, local_dim.tag): raise ValueError( f"The provided local dimension {local_dim} does not match any offset provider type." ) @@ -839,7 +839,7 @@ def _make_array_shape_and_strides( for dim in dims: if dim.kind == gtx_common.DimensionKind.LOCAL: # for local dimension, the size is taken from the associated connectivity type - shape.append(neighbor_table_types[dim.value].max_neighbors) + shape.append(neighbor_table_types[dim.tag].max_neighbors) elif gtx_dace_args.is_connectivity_identifier(name, self.offset_provider_type): # we use symbolic size for the global dimension of a connectivity shape.append(gtx_dace_args.field_size_symbol(name, dim, neighbor_table_types)) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py index ed6b8dcc9c..737648b911 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py @@ -254,7 +254,7 @@ 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.value) + offset_provider_type = sdfg_builder.get_offset_provider_type(local_dim.tag) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) local_idx = gtx_common.order_dimensions([*output_dims, local_dim]).index(local_dim) output_shape.insert(local_idx, offset_provider_type.max_neighbors) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index d05dcb8eaa..5e8a23d277 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 @@ -655,7 +655,7 @@ def _visit_deref(self, node: gtir.FunCall) -> DataExpr: assert len(field_desc.shape) == len(arg_expr.field_domain) field_indices = [(dim, arg_expr.indices[dim]) for dim, _ in arg_expr.field_domain] index_connectors = [ - IndexConnectorFmt.format(dim=dim.value) + IndexConnectorFmt.format(dim=gtx_common.codegen_name(dim.tag)) for dim, index in field_indices if not isinstance(index, SymbolExpr) ] @@ -665,7 +665,7 @@ def _visit_deref(self, node: gtir.FunCall) -> DataExpr: index_internals = ",".join( str(index.value - offset) if isinstance(index := arg_expr.indices[dim], SymbolExpr) - else f"{IndexConnectorFmt.format(dim=dim.value)} - {offset}" + else f"{IndexConnectorFmt.format(dim=dim.tag)} - {offset}" for (dim, offset) in arg_expr.field_domain ) deref_node, connector_mapping = self._add_tasklet( @@ -684,7 +684,7 @@ def _visit_deref(self, node: gtir.FunCall) -> DataExpr: # add termination points for the dynamic iterator indices for dim, index_expr in field_indices: - index_connector = IndexConnectorFmt.format(dim=dim.value) + index_connector = IndexConnectorFmt.format(dim=dim.tag) if isinstance(index_expr, MemletExpr): self._add_input_data_edge( index_expr.dc_node, @@ -769,7 +769,7 @@ 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.value), + self.subgraph_builder.get_offset_provider_type(local_dim.tag), gtx_common.NeighborConnectivityType, ) # find position of the local dimension in the field layout @@ -1155,7 +1155,11 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: self.sdfg, (conn_type.max_neighbors,), field_desc.dtype ) neighbors_node = self.state.add_access(neighbors_temp) - offset_type = gtx_common.Dimension(offset, gtx_common.DimensionKind.LOCAL) + # NOTE: the connectivity's own local dimension, not one synthesized from the offset + # tag. The latter named a local dimension after the *offset*, which only coincided with + # the real one under the old `V2EDim = Dimension("V2E")` convention, and under nominal + # identity (ADR 0028) a tag string cannot be turned back into a dimension at all. + offset_type = conn_type.neighbor_dim neighbor_idx = gtir_to_sdfg_utils.get_map_variable(offset_type) index_connector = "__index" @@ -1317,7 +1321,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if offset_type == _CONST_DIM: # this input argument is the result of `make_const_list` continue - offset_provider_t = self.subgraph_builder.get_offset_provider_type(offset_type.value) + offset_provider_t = self.subgraph_builder.get_offset_provider_type(offset_type.tag) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) input_conn_types[offset_type] = offset_provider_t @@ -1371,7 +1375,7 @@ 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.value) + conn_data = gtx_dace_args.connectivity_identifier(offset_type.tag) conn_desc = self.sdfg.arrays[conn_data] conn_desc.transient = False @@ -1433,7 +1437,7 @@ def _broadcast_const_list( ) -> ValueExpr: assert list_type.offset_type is not None offset_provider_t = self.subgraph_builder.get_offset_provider_type( - list_type.offset_type.value + list_type.offset_type.tag ) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) local_size = offset_provider_t.max_neighbors @@ -1473,7 +1477,7 @@ 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.value) + offset_provider_type = self.subgraph_builder.get_offset_provider_type(offset_type.tag) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) inp_conn = "_in" @@ -1485,7 +1489,7 @@ 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.value) + connectivity = gtx_dace_args.connectivity_identifier(offset_type.tag) self.sdfg.arrays[connectivity].transient = False reduce_node = gtx_library_nodes.ReduceWithSkipValues( 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 8e14ae41bd..ee4d56f320 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 @@ -326,7 +326,7 @@ def _construct_if_branch_output( 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.value + out_type.dtype.offset_type.tag ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) shape = [*shape, offset_provider_type.max_neighbors] diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index ba50b21dee..18b48987cb 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -384,7 +384,7 @@ 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.value) + offset_provider_type = sdfg_builder.get_offset_provider_type(offset_type.tag) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) list_size = offset_provider_type.max_neighbors return [scan_column_size, dace.symbolic.SymExpr(list_size)] diff --git a/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 aabd0f64ed..a2117a87a2 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 @@ -51,7 +51,7 @@ def get_map_variable(dim: gtx_common.Dimension) -> str: # 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_{dim.value}_gtx_{dim.kind}{suffix}" + return f"i_{gtx_common.codegen_name(dim.tag)}_gtx_{dim.kind}{suffix}" 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 f33fbf8bb5..4b3158c21b 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -77,7 +77,7 @@ def _field_symbol( offset_provider_type: gtx_common.OffsetProviderType | None, ) -> dace.symbol: if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is None: - name = f"__{field_name}_{dim.value}_{sym}" + name = f"__{field_name}_{gtx_common.codegen_name(dim.tag)}_{sym}" else: # a connectivity field assert offset_provider_type is not None assert m[1] in offset_provider_type @@ -109,22 +109,21 @@ def field_stride_symbol( return _field_symbol(field_name, dim, "stride", offset_provider_type) -def _range_symbol_name(field_name: str, axis: str) -> str: +def _range_symbol_name(field_name: str, dim: gtx_common.Dimension) -> str: """Common part of the name for the range start/stop symbols.""" - dim = gtx_common.Dimension(axis) field_range = im.call("get_domain_range")(field_name, dim) return gtir_python_codegen.get_source(field_range) def range_start_symbol(field_name: str, dim: gtx_common.Dimension) -> dace.symbol: """Format name of the start symbol for domain range.""" - name = f"{_range_symbol_name(field_name, dim.value)}_0" + name = f"{_range_symbol_name(field_name, dim)}_0" return dace.symbol(name, FIELD_SYMBOL_DTYPE) def range_stop_symbol(field_name: str, dim: gtx_common.Dimension) -> dace.symbol: """Format name of the stop symbol for domain range.""" - name = f"{_range_symbol_name(field_name, dim.value)}_1" + name = f"{_range_symbol_name(field_name, dim)}_1" return dace.symbol(name, FIELD_SYMBOL_DTYPE) diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py b/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py index 3e1a38f335..47b644a2b8 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py @@ -111,7 +111,7 @@ def __init__( super().__init__() if blocking_parameters is not None: self.blocking_parameters = [ - gtx_dace_lowering.get_map_variable(p) if isinstance(p, gtx_common.Dimension) else p + gtx_dace_lowering.get_map_variable(p) if isinstance(p, gtx_common.DimensionMeta) else p for p in blocking_parameters ] else: diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/map_orderer.py b/src/gt4py/next/program_processors/runners/dace/transformations/map_orderer.py index 2f56671da8..820680a17c 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/map_orderer.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/map_orderer.py @@ -125,7 +125,7 @@ def __init__( if unit_strides_dims is not None and unit_strides_kind is not None: raise ValueError("Specified both 'unit_strides_dims' and 'unit_strides_kind'.") elif unit_strides_dims is not None: - if isinstance(unit_strides_dims, (gtx_common.Dimension, str)): + if isinstance(unit_strides_dims, (gtx_common.DimensionMeta, str)): unit_strides_dims = [unit_strides_dims] self.unit_strides_dims = [ unit_strides_dim diff --git a/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 09f173d3f9..9bb62bd5de 100644 --- a/src/gt4py/next/program_processors/runners/roundtrip.py +++ b/src/gt4py/next/program_processors/runners/roundtrip.py @@ -60,11 +60,20 @@ def visit_Literal(self, node: itir.Literal, **kwargs: Any) -> str: return f"np.{dtype}(np.nan)" return node.value - OffsetLiteral = as_fmt("{value}") - AxisLiteral = as_fmt("{value}") + # NOTE: a tag is a qualified Python name, which is not a valid Python *identifier*, so the + # emitted program refers to each axis and offset through its mangled name. The header binds + # that name to the real object, looked up by the unmangled tag. + def visit_OffsetLiteral(self, node: itir.OffsetLiteral, **kwargs: Any) -> str: + # an integer offset literal is a shift amount, not a name + return common.codegen_name(node.value) if isinstance(node.value, str) else str(node.value) + + def visit_AxisLiteral(self, node: itir.AxisLiteral, **kwargs: Any) -> str: + return common.codegen_name(node.value) def visit_CartesianOffset(self, node: itir.CartesianOffset, **kwargs: Any) -> str: - return f"gtx.CartesianConnectivity({node.domain.value}, codomain={node.codomain.value})" + domain = common.codegen_name(node.domain.value) + codomain = common.codegen_name(node.codomain.value) + return f"gtx.CartesianConnectivity({domain}, codomain={codomain})" FunCall = as_fmt("{fun}({','.join(args)})") Lambda = as_mako("(lambda ${','.join(params)}: ${expr})") @@ -172,10 +181,13 @@ def _generate_source( """ ) - offset_literals_src = "\n".join(f'{o} = offset("{o}")' for o in offset_literals) + offset_literals_src = "\n".join( + f'{common.codegen_name(o)} = offset("{o}")' for o in offset_literals + ) + # A dimension is not constructed from its name any more: its tag is its qualified Python + # name, so the emitted program imports it (ADR 0028). axis_literals_src = "\n".join( - f'{o.value} = gtx.Dimension("{o.value}", kind=gtx.DimensionKind("{o.kind}"))' - for o in axis_literals_set + f'{common.codegen_name(o.value)} = gtx.resolve("{o.value}")' for o in axis_literals_set ) source_code = f"{header}{offset_literals_src}\n{axis_literals_src}\n{program}" diff --git a/src/gt4py/next/type_system/mypy_plugin.py b/src/gt4py/next/type_system/mypy_plugin.py index c3af362960..d1c423e325 100644 --- a/src/gt4py/next/type_system/mypy_plugin.py +++ b/src/gt4py/next/type_system/mypy_plugin.py @@ -19,11 +19,6 @@ The goal of this plugin is to reduce the amount of false positives from mypy that arise from correct usage of GT4Py. The following are examples for such false positives: -Dimensions are not fields: - - IDim = gtx.Dimension("IDim") - gtx.Field[gtx.Dims[IDim], float] # IDim is not a valid type - mixed precision math / different ways of describing the same dtype: a: gtx.Field[gtx.Dims[IDim], gtx.float64] @@ -35,10 +30,10 @@ The documentation on mypy plugins is at https://mypy.readthedocs.io/en/latest/extending_mypy.html -Known limitation: the fallback hook that rewrites stray 'Dimension' instances in type aliases -matches on the *variable name* ending in 'Dim' ('fullname.endswith("Dim")'). A dimension bound to -a name that does not end in 'Dim' -- e.g. 'I = gtx.Dimension("I")' -- silently gets no plugin -support at all. +Dimensions no longer need plugin support: a concrete dimension is a class +('class IDim(gtx.DimensionIndex): ...'), which is a valid annotation for any type checker. See ADR +0028. Only the mixed-precision hooks below remain; this plugin is scheduled for removal once +dtype-generic fields land. """ from __future__ import annotations @@ -46,22 +41,10 @@ import typing -def iter_dim_names() -> typing.Iterator[str]: - """Go through the four distinct place holders, then yield _AnyDim for everything after.""" - yield "_DimA" - yield "_DimB" - yield "_DimC" - yield "_DimD" - while True: - yield "_AnyDim" - - # if we can not import mypy we are not type checking with mypy, so we can skip all this try: from mypy import plugin as mplugin, types - DIM_MAP: dict[str, types.Type] = {} - FLOAT_TYPES = ["builtins.float", "numpy.float32", "numpy.float64"] INT_TYPES = [ "builtins.int", @@ -71,60 +54,6 @@ def iter_dim_names() -> typing.Iterator[str]: "numpy.signedinteger", ] - def fixup_dims_type(ctx: mplugin.AnalyzeTypeContext) -> types.Type: - """ - Overwrite Dims[...] to contain actual types. - - Example: - - CellDim = Dimension("CellDim") # this is not a type - a: Field[Dims[CellDim], gtx.float] # we transform this to Field[Dims[_DimA]] where _DimA *is* a type - - The actual types are defined in 'gt4py.next.common', if typing.TYPE_CHECKING is true. - """ - module_name = "gt4py.next.common" - dims = iter_dim_names() - args = [] - if ctx.type.args: - # replacement 'Dims[OneDim, OtherDim]' -> 'Dims[_DimA, _DimB]' happens here - for arg in ctx.type.args: - argname = getattr(arg, "name", "unknown") - if argname not in DIM_MAP: - DIM_MAP[argname] = ctx.api.analyze_type( - ctx.api.named_type(f"{module_name}.{next(dims)}", []) - ) - args.append(DIM_MAP[argname]) - else: - # do not accidentally replace 'Dims' -> 'Dims[Any]' (the former matches any number of dims, the latter only one) - args = [types.UnpackType(typ=types.AnyType(types.TypeOfAny.explicit))] - result = ctx.api.analyze_type(ctx.api.named_type("gt4py.next.common.Dims", args)) - return result - - def fixup_dims_from_typealiases(ctx: mplugin.AnalyzeTypeContext) -> types.Type: - """ - Catch 'Dimension' instances that made it through other replacements, this seems to happen in type aliases. - - Example: - - T = TypeVar("T", bound=(float, float64)) - CellDim = Dimension("CellDim") - CellField: TypeAlias = Field[Dims[CellDim], T] - - If we have seen this dimension instance before, reuse the same dim type in the replacement, else use `_AnyDim` - """ - result: types.Type | types.AnyType = types.AnyType( - types.TypeOfAny.explicit - ) # Fallback to Any if _AnyDim is not found - try: - if ctx.type.name not in DIM_MAP: - DIM_MAP[ctx.type.name] = ctx.api.analyze_type( - ctx.api.named_type("gt4py.next.common._AnyDim", []) - ) - result = DIM_MAP[ctx.type.name] - except AssertionError: # this probably happens when a dim type is analyzed in a context from where _AnyDim is unreachable - pass - return result - def blur_float_precision(ctx: mplugin.AnalyzeTypeContext) -> types.Type: """Turn everything into 'builtins.float'.""" # Note(ricoh): have tried to return numpy dtypes from here but ran into some error from mypy @@ -134,7 +63,7 @@ def blur_int_precision(ctx: mplugin.AnalyzeTypeContext) -> types.Type: """Turn everything into 'builtins.int'""" return ctx.api.named_type("builtins.int", []) - class TreatDimensionsAsTypes(mplugin.Plugin): + class BlurScalarPrecision(mplugin.Plugin): def get_type_analyze_hook( self, fullname: str ) -> typing.Callable[[mplugin.AnalyzeTypeContext], types.Type] | None: @@ -143,25 +72,22 @@ def get_type_analyze_hook( If a callback is returned, it has to return a mypy-representation of a valid type. """ - # replace Dims args with actual types - if fullname == "gt4py.next.common.Dims": - return fixup_dims_type - # replace stray 'Dimension' instances - elif fullname.endswith("Dim") and not fullname == "gt4py.next.common._AnyDim": - return fixup_dims_from_typealiases # treat all float precision types the same (GT4Py dsl will catch actual problems) - elif fullname in FLOAT_TYPES: + if fullname in FLOAT_TYPES: return blur_float_precision # treat all int precision types the same (GT4Py dsl will catch actual problems) elif fullname in INT_TYPES: return blur_int_precision return None + #: Deprecated alias: downstream configs may still name the old symbol. + TreatDimensionsAsTypes = BlurScalarPrecision + def plugin(version: str) -> type[mplugin.Plugin]: """ This is the entry point mypy looks for if this module was pointed to in config as a plugin. """ - return TreatDimensionsAsTypes + return BlurScalarPrecision except ImportError: pass diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 5e727821a5..0527d019df 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -903,9 +903,9 @@ def function_signature_incompatibilities_field( if field_type.dims and source_dim not in field_type.dims: yield ( f"Incompatible offset can not shift field defined on " - f"{', '.join([dim.value for dim in field_type.dims])} from " - f"{source_dim.value} to target dim(s): " - f"{', '.join([dim.value for dim in target_dims])}" + 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])}" ) diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index ca5cf81a3a..a2bfc86831 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -122,7 +122,7 @@ class FieldType(DataType, CallableType): dtype: ScalarType | ListType def __str__(self) -> str: - dims = "..." if self.dims is Ellipsis else f"[{', '.join(dim.value for dim in self.dims)}]" + dims = "..." if self.dims is Ellipsis else f"[{', '.join(dim.__qualname__ for dim in self.dims)}]" return f"Field[{dims}, {self.dtype}]" @eve_datamodels.validator("dims") diff --git a/src/gt4py/next/type_system/type_translation.py b/src/gt4py/next/type_system/type_translation.py index df721da0e7..aa7730b97d 100644 --- a/src/gt4py/next/type_system/type_translation.py +++ b/src/gt4py/next/type_system/type_translation.py @@ -219,7 +219,7 @@ def from_type_hint( ) if isinstance(dim_arg, list): for d in dim_arg: - if not isinstance(d, common.Dimension): + if not isinstance(d, common.DimensionMeta): raise ValueError(f"Invalid field dimension definition '{d}'.") dims.append(d) else: @@ -342,7 +342,7 @@ def from_value(value: Any) -> ts.TypeSpec: f"Value '{value}' is out of range to be representable as 'INT32' or 'INT64'." ) return candidate_type - elif isinstance(value, common.Dimension): + elif isinstance(value, common.DimensionMeta): symbol_type = ts.DimensionType(dim=value) elif isinstance(value, common.Field): dims = list(value.domain.dims) diff --git a/tests/next_tests/artifacts/custom_named_collections.py b/tests/next_tests/artifacts/custom_named_collections.py index d7e9cb1bf1..0007cc12f3 100644 --- a/tests/next_tests/artifacts/custom_named_collections.py +++ b/tests/next_tests/artifacts/custom_named_collections.py @@ -15,11 +15,20 @@ from gt4py import next as gtx from gt4py.eve.xtyping import NestedTuple -from gt4py.next import common, Dimension, Field, float32, float64, Dims, named_collections +from gt4py.next import ( + common, + Dimension, + DimensionIndex, + Field, + float32, + float64, + Dims, + named_collections, +) from gt4py.next.type_system import type_specifications as ts -TDim = Dimension("TDim") # Meaningless dimension just for tests +class TDim(DimensionIndex): ... class SingleElementNamedTupleNamedCollection(NamedTuple): diff --git a/tests/next_tests/benchmarks/benchmark_program_call.py b/tests/next_tests/benchmarks/benchmark_program_call.py index 166509e8f3..9f5f90b992 100644 --- a/tests/next_tests/benchmarks/benchmark_program_call.py +++ b/tests/next_tests/benchmarks/benchmark_program_call.py @@ -41,8 +41,10 @@ from pytest_benchmark import fixture as ptb_fixture -Cell = gtx.Dimension("Cell") -IDim = gtx.Dimension("IDim") +class Cell(gtx.DimensionIndex): ... + + +class IDim(gtx.DimensionIndex): ... @pytest.mark.parametrize("backend", BACKENDS, ids=lambda b: b.name) diff --git a/tests/next_tests/fixtures/past_common.py b/tests/next_tests/fixtures/past_common.py index 0b2681059c..83c13988fe 100644 --- a/tests/next_tests/fixtures/past_common.py +++ b/tests/next_tests/fixtures/past_common.py @@ -14,8 +14,10 @@ from gt4py.next import float64 -IDim = gtx.Dimension("IDim") -JDim = gtx.Dimension("JDim") +class IDim(gtx.DimensionIndex): ... + + +class JDim(gtx.DimensionIndex): ... # TODO(tehrengruber): Improve test structure. Identity needs to be decorated diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 18617a6f91..464bf7fa1f 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -31,6 +31,13 @@ import next_tests +# NOTE: the unstructured dimensions are declared once, in `toy_connectivity`, and imported here. +# Both modules used to declare their own `Dimension("Vertex")` etc., which compared equal; under +# nominal identity (ADR 0028) that would be two different dimensions, and tests that mix a +# `toy_connectivity` connectivity with a `cases_utils` mesh would silently stop matching. +from next_tests.toy_connectivity import C2EDim, Cell, E2VDim, Edge, V2EDim, Vertex + + __all__ = [ "exec_alloc_descriptor", "mesh_descriptor", @@ -152,29 +159,38 @@ def debug_itir(tree): DimsType = TypeVar("DimsType") DType = TypeVar("DType") -IDim = gtx.Dimension("IDim") + +class IDim(gtx.DimensionIndex): ... + + IHalfDim = common.flip_staggered(IDim) -JDim = gtx.Dimension("JDim") + + +class JDim(gtx.DimensionIndex): ... + + JHalfDim = common.flip_staggered(JDim) -KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + + +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + KHalfDim = common.flip_staggered(KDim) Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) -Vertex = gtx.Dimension("Vertex") -Edge = gtx.Dimension("Edge") -Cell = gtx.Dimension("Cell") + EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) -V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) -E2VDim = gtx.Dimension("E2V", kind=gtx.DimensionKind.LOCAL) -C2EDim = gtx.Dimension("C2E", kind=gtx.DimensionKind.LOCAL) -C2VDim = gtx.Dimension("C2V", kind=gtx.DimensionKind.LOCAL) -V2E = gtx.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset("E2V", source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset("C2E", source=Edge, target=(Cell, C2EDim)) -C2V = gtx.FieldOffset("C2V", source=Vertex, target=(Cell, C2VDim)) + +class C2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) +E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) +C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2V = gtx.FieldOffset(C2VDim.tag, source=Vertex, target=(Cell, C2VDim)) size = 10 diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py index e984140038..7dc8943e81 100644 --- a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py +++ b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py @@ -108,7 +108,9 @@ def test_sdfgConvertible_connectivities(unstructured_case): # noqa: F811 allocator=allocator, ) - testee2 = testee.with_backend(backend).with_compilation_options(connectivities={"E2V": e2v}) + testee2 = testee.with_backend(backend).with_compilation_options( + connectivities={E2VDim.tag: e2v} + ) @dace.program def sdfg( @@ -122,7 +124,7 @@ def sdfg( ) return out - connectivities = {"E2V": e2v} # replace 'e2v' with 'e2v.__gt_type__()' when GTIR is AOT + connectivities = {E2VDim.tag: e2v} # replace 'e2v' with 'e2v.__gt_type__()' when GTIR is AOT offset_provider = OffsetProvider_t.dtype._typeclass.as_ctypes()(E2V=e2v.data_ptr()) a = gtx.as_field([Vertex], xp.asarray([0.0, 1.0, 2.0]), allocator=allocator) 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 e613e3db92..fa65a38b40 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 @@ -30,13 +30,22 @@ from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations -IDim = gtx.Dimension("I") +class IDim(gtx.DimensionIndex): ... + + I_SIZE = 8 -Cell = gtx.Dimension("Cell") -Edge = gtx.Dimension("Edge") -C2EDim = gtx.Dimension("C2E", kind=gtx.DimensionKind.LOCAL) -C2E = gtx.FieldOffset("C2E", source=Edge, target=(Cell, C2EDim)) + +class Cell(gtx.DimensionIndex): ... + + +class Edge(gtx.DimensionIndex): ... + + +class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) C2E_TABLE = np.array( [ @@ -213,7 +222,7 @@ def test_write_back_buffer_elimination_from_lowering_with_reduction( survived until the transformation runs. """ offset_provider = { - "C2E": constructors.as_connectivity( + C2EDim.tag: constructors.as_connectivity( domain={Cell: N_CELLS, C2EDim: 4}, codomain=Edge, data=C2E_TABLE, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py index 0ed7ed60f0..dfe013e45a 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py @@ -454,7 +454,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: t = concat_where(Vertex < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -477,7 +477,7 @@ def testee( t = concat_where(KDim < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case_3d.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() k_mask = np.arange(unstructured_case_3d.default_sizes[KDim]) < 2 cases.verify_with_default_data( unstructured_case_3d, @@ -497,7 +497,7 @@ def test_with_local_and_nonlocal_field(unstructured_case, static_domains: bool): def testee(a: cases.EField, b: cases.VField) -> cases.VField: return neighbor_sum(concat_where(Vertex < 2, a(V2E), b), axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -524,7 +524,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), b(V2E)), (c(V2E), d(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -557,7 +557,7 @@ def testee(a: cases.EField) -> tuple[cases.VField, cases.VField]: neighbor_sum(concat_where(Vertex < 2, 3, a(V2E)), axis=V2EDim), ) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -590,7 +590,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), c), (3, b(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py index 0fc370adbb..753015f0d7 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py @@ -31,11 +31,11 @@ def testee( ) # multiplication with shifted `ones` because reduction of only non-shifted field with local dimension is not supported inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider["V2E"].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() ) ones = cases.allocate(unstructured_case, testee, "ones").strategy(cases.ConstInitializer(1))() - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() cases.verify( unstructured_case, testee, @@ -55,7 +55,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return inp[V2EDim(0)] + inp[V2EDim(1)] + inp[V2EDim(2)] + inp[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider["V2E"].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() ) cases.verify( @@ -77,7 +77,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int64 return inp_64[V2EDim(0)] + inp_64[V2EDim(1)] + inp_64[V2EDim(2)] + inp_64[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider["V2E"].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() ) cases.verify( @@ -99,7 +99,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return neighbor_sum(inp, axis=V2EDim) inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider["V2E"].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() ) cases.verify( @@ -107,7 +107,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 testee, inp, out=cases.allocate(unstructured_case, testee, cases.RETURN)(), - ref=np.sum(unstructured_case.offset_provider["V2E"].asnumpy(), axis=1), + ref=np.sum(unstructured_case.offset_provider[V2EDim.tag].asnumpy(), axis=1), ) @@ -119,7 +119,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: return inp(V2E) out = unstructured_case.as_field( - [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider["V2E"].asnumpy()) + [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider[V2EDim.tag].asnumpy()) ) inp = cases.allocate(unstructured_case, testee, "inp")() cases.verify( @@ -127,5 +127,5 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: testee, inp, out=out, - ref=inp.asnumpy()[unstructured_case.offset_provider["V2E"].asnumpy()], + ref=inp.asnumpy()[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py index 8ba6f3091a..c5cbd43653 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py @@ -11,12 +11,27 @@ import pytest -from gt4py.next import Dimension, DimensionKind, Field, field_operator, int32, int64, scan_operator +from gt4py.next import ( + Dimension, + DimensionIndex, + DimensionKind, + Field, + field_operator, + int32, + int64, + scan_operator, +) from gt4py.next.ffront.ast_passes import single_static_assign as ssa from gt4py.next.ffront.foast_pretty_printer import pretty_format from gt4py.next.ffront.func_to_foast import FieldOperatorParser +class I(DimensionIndex): ... + + +class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... + + @pytest.mark.parametrize( "test_case", [ @@ -45,7 +60,6 @@ def test_one_to_one(test_case: str): def test_fieldop(): - I = Dimension("I") @field_operator def foo(inp1: Field[[I], int64], inp2: Field[[I], int64]): @@ -67,7 +81,6 @@ def bar(inp1: Field[[I], int64], inp2: Field[[I], int64]) -> Field[[I], int64]: def test_scanop(): - KDim = Dimension("KDim", kind=DimensionKind.VERTICAL) @scan_operator(axis=KDim, forward=False, init=1) def scan(inp: int32) -> int32: diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py index 911ace0804..39aee5de06 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py @@ -483,7 +483,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) ) @@ -641,7 +641,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py index 311b1870ed..dd0cc6fb43 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py @@ -17,6 +17,7 @@ from next_tests.integration_tests import cases from next_tests.integration_tests.cases import ( + C2EDim, C2E, E2V, V2E, @@ -52,7 +53,7 @@ def testee(edge_f: cases.EField) -> cases.VField: inp = cases.allocate(unstructured_case, testee, "edge_f", strategy=strategy)() out = cases.allocate(unstructured_case, testee, cases.RETURN)() - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() ref = np.max( inp.asnumpy()[v2e_table], axis=1, @@ -69,7 +70,7 @@ def minover(edge_f: cases.EField) -> cases.VField: out = min_over(edge_f(V2E), axis=V2EDim) return out - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() cases.verify_with_default_data( unstructured_case, minover, @@ -99,7 +100,7 @@ def reduction_ek_field( "fop", [reduction_e_field, reduction_ek_field], ids=lambda fop: fop.__name__ ) def test_neighbor_sum(unstructured_case_3d, fop): - v2e_table = unstructured_case_3d.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() edge_f = cases.allocate(unstructured_case_3d, fop, "edge_f")() @@ -151,7 +152,7 @@ def fencil_op(edge_f: EKField) -> VKField: def fencil(edge_f: EKField, out: VKField): fencil_op(edge_f, out=out) - v2e_table = unstructured_case_3d.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() field = cases.allocate(unstructured_case_3d, fencil, "edge_f", sizes={KDim: 2})() out = cases.allocate(unstructured_case_3d, fencil_op, cases.RETURN, sizes={KDim: 1})() @@ -184,7 +185,7 @@ def reduce_expr(edge_f: cases.EField) -> cases.VField: def fencil(edge_f: cases.EField, out: cases.VField): reduce_expr(edge_f, out=out) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() cases.verify_with_default_data( unstructured_case, fencil, @@ -206,7 +207,7 @@ def test_reduction_with_common_expression(unstructured_case): def testee(flux: cases.EField) -> cases.VField: return neighbor_sum(flux(V2E) + flux(V2E), axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -222,7 +223,7 @@ def test_reduction_expression_with_where(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, inp(V2E), inp(V2E)), axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -251,7 +252,7 @@ def test_reduction_expression_with_where_and_tuples(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, (inp(V2E), inp(V2E)), (inp(V2E), inp(V2E)))[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -280,7 +281,7 @@ def test_reduction_expression_with_where_and_scalar(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(inp(V2E) + where(mask, inp(V2E), 1), axis=V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -325,7 +326,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider["E2V"].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], ) @@ -346,7 +347,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() neighbor_0_iter = iter(enumerate(e2v_table[:, 0])) edge_start = next(i for i, v in neighbor_0_iter if v >= ORIGIN) edge_stop = next(i for i, v in neighbor_0_iter if v < ORIGIN) @@ -391,16 +392,16 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_flat, - ref=lambda inp: inp[unstructured_case.offset_provider["E2V"].asnumpy()[:, 0]][ - unstructured_case.offset_provider["C2E"].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] ], ) cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_intermediate_result, - ref=lambda inp: inp[unstructured_case.offset_provider["E2V"].asnumpy()[:, 0]][ - unstructured_case.offset_provider["C2E"].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] ], comparison=lambda inp, tmp: np.all(inp == tmp), ) @@ -408,8 +409,8 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured, - ref=lambda inp: inp[unstructured_case.offset_provider["E2V"].asnumpy()[:, 0]][ - unstructured_case.offset_provider["C2E"].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] ], ) @@ -431,7 +432,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() neighbor_iter = iter(enumerate(e2v_table)) edge_start = next(i for i, v in neighbor_iter if all(v >= ORIGIN)) edge_stop = next(i for i, v in neighbor_iter if any(v < ORIGIN)) @@ -452,11 +453,12 @@ def testee(a: cases.VField) -> cases.VField: unstructured_case, testee, ref=lambda a: np.sum( - np.sum(a[unstructured_case.offset_provider["E2V"].asnumpy()], axis=1, initial=0)[ - unstructured_case.offset_provider["V2E"].asnumpy() + np.sum(a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()], axis=1, initial=0)[ + unstructured_case.offset_provider[V2EDim.tag].asnumpy() ], axis=1, - where=unstructured_case.offset_provider["V2E"].asnumpy() != common._DEFAULT_SKIP_VALUE, + where=unstructured_case.offset_provider[V2EDim.tag].asnumpy() + != common._DEFAULT_SKIP_VALUE, ), comparison=lambda a, tmp_2: np.all(a == tmp_2), ) @@ -477,8 +479,8 @@ def testee(inp: cases.EField) -> cases.EField: unstructured_case, testee, ref=lambda inp: np.sum( - np.sum(inp[unstructured_case.offset_provider["V2E"].asnumpy()], axis=1)[ - unstructured_case.offset_provider["E2V"].asnumpy() + np.sum(inp[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1)[ + unstructured_case.offset_provider[E2VDim.tag].asnumpy() ], axis=1, ), @@ -495,7 +497,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: tmp = red(E2V[0]) return tmp - v2e = unstructured_case.offset_provider["V2E"] + v2e = unstructured_case.offset_provider[V2EDim.tag] cases.verify_with_default_data( unstructured_case, reduce_tuple_element, @@ -504,7 +506,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: axis=1, initial=0, where=v2e.asnumpy() != common._DEFAULT_SKIP_VALUE, - )[unstructured_case.offset_provider["E2V"].asnumpy()[:, 0]], + )[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], ) @@ -516,7 +518,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: tmp = neighbor_sum(b(V2E) if 2 < 3 else a(V2E), axis=V2EDim) return tmp - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -538,7 +540,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex], int32]: inp = cases.allocate(unstructured_case, testee, "inp")() - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() cases.verify( unstructured_case, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py index 5b7430fc91..ac348266f3 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py @@ -13,6 +13,7 @@ from next_tests.integration_tests import cases from next_tests.integration_tests.cases import ( + E2VDim, E2V, Edge, IDim, @@ -222,7 +223,7 @@ def testee( )() out = cases.allocate(unstructured_case_3d, testee, cases.RETURN)() - e2v_table = unstructured_case_3d.offset_provider["E2V"].asnumpy() + e2v_table = unstructured_case_3d.offset_provider[E2VDim.tag].asnumpy() cases.verify( unstructured_case_3d, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py index 8900154140..3db69adf0e 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py @@ -17,6 +17,7 @@ from next_tests.integration_tests import cases from next_tests.integration_tests.cases import ( + E2VDim, E2V, Case, KDim, @@ -38,9 +39,9 @@ def exec_alloc_descriptor(): translation=functools.partial( gtfn.make_gtfn_translation, symbolic_domain_sizes={ - "Cell": "num_cells", - "Edge": "num_edges", - "Vertex": "num_vertices", + Cell.tag: "num_cells", + Edge.tag: "num_edges", + Vertex.tag: "num_vertices", }, ), ) @@ -81,7 +82,9 @@ def test_verification(testee, exec_alloc_descriptor, mesh_descriptor): a = cases.allocate(unstructured_case, testee, "a")() out = cases.allocate(unstructured_case, testee, "out")() - first_nbs, second_nbs = (mesh_descriptor.offset_provider["E2V"].asnumpy()[:, i] for i in [0, 1]) + first_nbs, second_nbs = ( + mesh_descriptor.offset_provider[E2VDim.tag].asnumpy()[:, i] for i in [0, 1] + ) ref = (a.ndarray * 2)[first_nbs] + (a.ndarray * 2)[second_nbs] cases.verify( diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index 91025f156f..c55a145314 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -164,8 +164,8 @@ def testee(a: cases.EField, b: cases.EField) -> tuple[cases.VField, cases.VField unstructured_case, testee, ref=lambda a, b: [ - np.sum(a[unstructured_case.offset_provider["V2E"].asnumpy()], axis=1), - np.sum(b[unstructured_case.offset_provider["V2E"].asnumpy()], axis=1), + np.sum(a[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), + np.sum(b[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), ], comparison=lambda a, tmp: (np.all(a[0] == tmp[0]), np.all(a[1] == tmp[1])), ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py index 2c12d74f50..6a160cefc3 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py @@ -51,7 +51,7 @@ def testee(a: gtx.Field[[Vertex], np.float64]) -> gtx.Field[[Edge], int64]: tmp = astype(a(E2V), int64) return neighbor_sum(tmp, axis=E2VDim) - e2v_table = unstructured_case.offset_provider["E2V"].asnumpy() + e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() cases.verify_with_default_data( unstructured_case, 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 9f51689d77..3700eac616 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 @@ -26,7 +26,7 @@ BACKENDS = [None, gtfn_cpu] -IDim = gtx.Dimension("IDim") +class IDim(gtx.DimensionIndex): ... @gtx.field_operator 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 b60626acd9..b16461447a 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 @@ -56,6 +56,12 @@ from next_tests.unit_tests.conftest import program_processor, run_processor +class Node(gtx.DimensionIndex): ... + + +class NeighDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + def array_maker(*lists): def _listify(val): if isinstance(val, Iterable): @@ -67,7 +73,7 @@ def _listify(val): return res -IDim = gtx.Dimension("IDim") +class IDim(gtx.DimensionIndex): ... def field_maker(*arrays): @@ -245,9 +251,6 @@ def foo(a): def test_can_deref(program_processor, stencil): program_processor, validate = program_processor - Node = gtx.Dimension("Node") - NeighDim = gtx.Dimension("Neighbor", kind=gtx.DimensionKind.LOCAL) - inp = gtx.as_field([Node], np.ones((1,), dtype=np.int32)) out = gtx.as_field([Node], np.asarray([0], dtype=inp.dtype)) diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py index eae66d425b..ed4039567d 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py @@ -16,7 +16,7 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -IDim = gtx.Dimension("IDim") +class IDim(gtx.DimensionIndex): ... @fundef diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py index 66a1edc189..ed1ed9d6a8 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py @@ -16,7 +16,8 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -I = gtx.Dimension("I") +class I(gtx.DimensionIndex): ... + _isize = 10 diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py index 5a386ee807..7220830aab 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py @@ -25,7 +25,9 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -I = gtx.Dimension("I") +class I(gtx.DimensionIndex): ... + + Ioff = gtx.CartesianConnectivity(I) 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 68e5f9d532..bd675b5f51 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py @@ -17,13 +17,21 @@ from gt4py.next.iterator.embedded import StridedConnectivityField -LocA = gtx.Dimension("LocA") -LocAB = gtx.Dimension("LocAB") -LocB = gtx.Dimension("LocB") # unused +class Dummy(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +class LocA(gtx.DimensionIndex): ... + + +class LocAB(gtx.DimensionIndex): ... + + +class LocB(gtx.DimensionIndex): ... + LocA2LocAB = offset("O") LocA2LocAB_offset_provider = StridedConnectivityField( - domain_dims=(LocA, gtx.Dimension("Dummy", kind=gtx.DimensionKind.LOCAL)), + domain_dims=(LocA, Dummy), codomain_dim=LocAB, max_neighbors=2, ) diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py index 95a1eb67f0..79fcaf1a83 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py @@ -16,9 +16,14 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -IDim = gtx.Dimension("IDim") -JDim = gtx.Dimension("JDim") -KDim = gtx.Dimension("KDim") +class IDim(gtx.DimensionIndex): ... + + +class JDim(gtx.DimensionIndex): ... + + +class KDim(gtx.DimensionIndex): ... + # semantics of stencil return that is called from the fencil (after `:` the structure of the output) # `return a` -> a: field diff --git a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_ffront_fvm_nabla.py b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_ffront_fvm_nabla.py index 2b4847d3c3..daaf9d05f8 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_ffront_fvm_nabla.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_ffront_fvm_nabla.py @@ -88,7 +88,7 @@ def test_ffront_compute_zavgS(exec_alloc_descriptor): setup.input_field, setup.S_fields[0], out=zavgS, - offset_provider={"E2V": setup.edges2node_connectivity}, + offset_provider={E2VDim.tag: setup.edges2node_connectivity}, ) assert_close(-199755464.25741270, np.min(zavgS.asnumpy())) @@ -113,8 +113,8 @@ def test_ffront_nabla(exec_alloc_descriptor): setup.vol_field, out=(pnabla_MXX, pnabla_MYY), offset_provider={ - "E2V": setup.edges2node_connectivity, - "V2E": setup.nodes2edge_connectivity, + E2VDim.tag: setup.edges2node_connectivity, + V2EDim.tag: setup.nodes2edge_connectivity, }, ) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_icon_like_scan.py b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_icon_like_scan.py index 234b8f4d02..861ef70526 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_icon_like_scan.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_icon_like_scan.py @@ -30,10 +30,6 @@ ] -Cell = gtx.Dimension("Cell") -KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) - - class State(NamedTuple): z_q_new: float w_new: float diff --git a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py index 0ddb42a521..01e4ed92bb 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py @@ -14,6 +14,8 @@ import gt4py.next as gtx from next_tests.integration_tests import cases from next_tests.integration_tests.cases import ( + E2VDim, + C2EDim, IDim, JDim, KDim, @@ -544,7 +546,7 @@ def test_program_unstructured(unstructured_case): unstructured_case.default_sizes[Cell], unstructured_case.default_sizes[Edge], inout=(out_a_shifted, out_a), - ref=((a.ndarray)[unstructured_case.offset_provider["C2E"].asnumpy()[:, 1]], a), + ref=((a.ndarray)[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]], a), ) @@ -598,7 +600,7 @@ def test_program_temporary(unstructured_case): extend={Cell: (-restrict_cell[0], restrict_cell[1])}, )() - e2v = (a.ndarray)[unstructured_case.offset_provider["E2V"].asnumpy()[:, 1]] + e2v = (a.ndarray)[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 1]] cases.verify( unstructured_case, prog_temporary, @@ -614,7 +616,7 @@ def test_program_temporary(unstructured_case): inout=(out_edge, out_cell), ref=( e2v[restrict_edge[0] : edge_size + restrict_edge[1]], - e2v[unstructured_case.offset_provider["C2E"].asnumpy()[:, 1]][ + e2v[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]][ restrict_cell[0] : cell_size + restrict_cell[1] ], ), diff --git a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py index 84c8937b7e..8c3a8095b8 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py @@ -36,13 +36,20 @@ from gt4py.next.iterator import atlas_utils -Vertex = gtx.Dimension("Vertex") -Edge = gtx.Dimension("Edge") -V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) -E2VDim = gtx.Dimension("E2V", kind=gtx.DimensionKind.LOCAL) +class Vertex(gtx.DimensionIndex): ... -V2E = gtx.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset("E2V", source=Vertex, target=(Edge, E2VDim)) + +class Edge(gtx.DimensionIndex): ... + + +class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) +E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) def assert_close(expected, actual): diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py index 6ae2af59a9..b370bf7810 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py @@ -24,9 +24,13 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -IDim = gtx.Dimension("IDim") -JDim = gtx.Dimension("JDim") -KDim = gtx.Dimension("KDim") +class IDim(gtx.DimensionIndex): ... + + +class JDim(gtx.DimensionIndex): ... + + +class KDim(gtx.DimensionIndex): ... # cross-reference why new type inference does not support this diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py index e7984f7cd3..3697652641 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py @@ -30,19 +30,22 @@ ) from gt4py.next.iterator.runtime import set_at, fendef, fundef, offset +# NOTE: the dimensions are imported, not redeclared. The connectivities come from +# `nabla_setup` and are built on *its* dimension classes; under nominal identity (ADR 0028) a +# same-named redeclaration here would be a different dimension, where it used to compare equal. from next_tests.integration_tests.multi_feature_tests.fvm_nabla_setup import ( + E2VDim, + Edge, + V2EDim, + Vertex, assert_close, nabla_setup, ) from next_tests.unit_tests.conftest import program_processor, run_processor -Vertex = gtx.Dimension("Vertex") -Edge = gtx.Dimension("Edge") -V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) - -V2E = offset("V2E") -E2V = offset("E2V") +V2E = offset(V2EDim.tag) +E2V = offset(E2VDim.tag) @fundef @@ -119,7 +122,7 @@ def test_compute_zavgS(program_processor): zavgS, setup.input_field, setup.S_fields[0], - offset_provider={"E2V": setup.edges2node_connectivity}, + offset_provider={E2VDim.tag: setup.edges2node_connectivity}, ) if validate: @@ -133,7 +136,7 @@ def test_compute_zavgS(program_processor): zavgS, setup.input_field, setup.S_fields[1], - offset_provider={"E2V": setup.edges2node_connectivity}, + offset_provider={E2VDim.tag: setup.edges2node_connectivity}, ) if validate: assert_close(-1000788897.3202186, np.min(zavgS.asnumpy())) @@ -167,7 +170,7 @@ def test_compute_zavgS2(program_processor): zavgS, setup.input_field, setup.S_fields, - offset_provider={"E2V": setup.edges2node_connectivity}, + offset_provider={E2VDim.tag: setup.edges2node_connectivity}, ) if validate: @@ -203,8 +206,8 @@ def test_nabla(program_processor): setup.sign_field, setup.vol_field, offset_provider={ - "E2V": setup.edges2node_connectivity, - "V2E": setup.nodes2edge_connectivity, + E2VDim.tag: setup.edges2node_connectivity, + V2EDim.tag: setup.nodes2edge_connectivity, }, ) @@ -243,8 +246,8 @@ def test_nabla2(program_processor): setup.sign_field, setup.vol_field, offset_provider={ - "E2V": setup.edges2node_connectivity, - "V2E": setup.nodes2edge_connectivity, + E2VDim.tag: setup.edges2node_connectivity, + V2EDim.tag: setup.nodes2edge_connectivity, }, ) @@ -321,8 +324,8 @@ def test_nabla_sign(program_processor): vertex_index, # TODO(havogt): should be an index function field setup.is_pole_edge_field, offset_provider={ - "E2V": setup.edges2node_connectivity, - "V2E": setup.nodes2edge_connectivity, + E2VDim.tag: setup.edges2node_connectivity, + V2EDim.tag: setup.nodes2edge_connectivity, }, ) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py index 1aac9c76d1..746e9f83ed 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py @@ -22,7 +22,7 @@ def multiply(alpha, inp): return deref(alpha) * deref(inp) -IDim = gtx.Dimension("IDim") +class IDim(gtx.DimensionIndex): ... @pytest.mark.uses_ir_if_stmts diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py index 60aef83302..6db3cafd2d 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py @@ -24,8 +24,11 @@ from next_tests.unit_tests.conftest import program_processor_no_transforms, run_processor -IDim = gtx.Dimension("IDim") -JDim = gtx.Dimension("JDim") +class IDim(gtx.DimensionIndex): ... + + +class JDim(gtx.DimensionIndex): ... + i = gtx.CartesianConnectivity(IDim) j = gtx.CartesianConnectivity(JDim) 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 ad2337237d..fdf8cf5114 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 @@ -27,6 +27,7 @@ from gt4py.next.program_processors.runners import gtfn from next_tests.toy_connectivity import ( + C2EDim, C2E, E2V, V2E, @@ -93,7 +94,7 @@ def test_sum_edges_to_vertices(program_processor, stencil): program_processor, inp, out=out, - offset_provider={"V2E": v2e_conn}, + offset_provider={V2EDim.tag: v2e_conn}, ) if validate: assert np.allclose(out.asnumpy(), ref) @@ -115,7 +116,7 @@ def test_map_neighbors(program_processor): program_processor, inp, out=out, - offset_provider={"V2E": v2e_conn}, + offset_provider={V2EDim.tag: v2e_conn}, ) if validate: assert np.allclose(out.asnumpy(), ref) @@ -138,7 +139,7 @@ def test_map_make_const_list(program_processor): program_processor, inp, out=out, - offset_provider={"V2E": v2e_conn}, + offset_provider={V2EDim.tag: v2e_conn}, ) if validate: assert np.allclose(out.asnumpy(), ref) @@ -162,8 +163,8 @@ def test_first_vertex_neigh_of_first_edge_neigh_of_cells_fencil(program_processo inp, out=out, offset_provider={ - "E2V": e2v_conn, - "C2E": c2e_conn, + E2VDim.tag: e2v_conn, + C2EDim.tag: c2e_conn, }, ) if validate: @@ -191,7 +192,7 @@ def test_sparse_input_field(program_processor): non_sparse, inp, out=out, - offset_provider={"V2E": v2e_conn}, + offset_provider={V2EDim.tag: v2e_conn}, ) if validate: @@ -215,8 +216,8 @@ def test_sparse_input_field_v2v(program_processor): inp, out=out, offset_provider={ - "V2V": v2v_conn, - "V2E": v2e_conn, + V2VDim.tag: v2v_conn, + V2EDim.tag: v2e_conn, }, ) @@ -242,7 +243,7 @@ def test_slice_sparse(program_processor): program_processor, inp, out=out, - offset_provider={"V2V": v2v_conn}, + offset_provider={V2VDim.tag: v2v_conn}, ) if validate: @@ -266,7 +267,7 @@ def test_slice_twice_sparse(program_processor): program_processor, inp, out=out, - offset_provider={"V2V": v2v_conn}, + offset_provider={V2VDim.tag: v2v_conn}, ) if validate: @@ -291,7 +292,7 @@ def test_slice_shifted_sparse(program_processor): program_processor, inp, out=out, - offset_provider={"V2V": v2v_conn}, + offset_provider={V2VDim.tag: v2v_conn}, ) if validate: @@ -320,7 +321,7 @@ def test_lift(program_processor): program_processor, inp, out=out, - offset_provider={"V2V": v2v_conn}, + offset_provider={V2VDim.tag: v2v_conn}, ) if validate: assert np.allclose(out.asnumpy(), ref) @@ -343,7 +344,7 @@ def test_shift_sparse_input_field(program_processor): program_processor, inp, out=out, - offset_provider={"V2V": v2v_conn}, + offset_provider={V2VDim.tag: v2v_conn}, ) if validate: @@ -373,8 +374,8 @@ def test_shift_sparse_input_field2(program_processor): out2 = gtx.as_field([Vertex], np.zeros([9], dtype=inp.dtype)) offset_provider = { - "E2V": e2v_conn, - "V2E": v2e_conn, + E2VDim.tag: e2v_conn, + V2EDim.tag: v2e_conn, } domain = {Vertex: range(0, 9)} @@ -428,7 +429,7 @@ def test_sparse_shifted_stencil_reduce(program_processor): program_processor, inp, out=out, - offset_provider={"V2V": v2v_conn}, + offset_provider={V2VDim.tag: v2v_conn}, ) if validate: diff --git a/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py b/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py index b69950928d..b5588ff7ed 100644 --- a/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py +++ b/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py @@ -10,8 +10,11 @@ from gt4py.next import common -I = common.Dimension("I") -J = common.Dimension("J") + +class I(common.DimensionIndex): ... + + +class J(common.DimensionIndex): ... def test_domain_pickle_after_slice(): 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 89a72b428f..689d0d5f71 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 @@ -36,15 +36,23 @@ ) -V = gtx.Dimension("V") -E = gtx.Dimension("E") +class V(gtx.DimensionIndex): ... + + +class E(gtx.DimensionIndex): ... + + +#: N1 == N3 == N4, but N2 differs: the tag is `TaggedOffDim.tag`, the variable is `off_a`. +class TaggedOffDim(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +off_a = gtx.FieldOffset(TaggedOffDim.tag, source=E, target=(V, TaggedOffDim)) -#: N1 == N3 == N4, but N2 differs: the tag is `TaggedOff`, the variable is `off_a`. -TaggedOffDim = gtx.Dimension("TaggedOff", kind=common.DimensionKind.LOCAL) -off_a = gtx.FieldOffset("TaggedOff", source=E, target=(V, TaggedOffDim)) #: N1 == N2 == N4, but N3 differs: the local dimension is `Neigh`, the tag is `OffB`. -Neigh = gtx.Dimension("Neigh", kind=common.DimensionKind.LOCAL) +class Neigh(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) @@ -61,7 +69,7 @@ def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) # NOTE: `.asnumpy()`, not `.ndarray`: under a GPU allocator the latter is a device # array, and `simple_mesh` builds the table from NumPy anyway. - v2e_arr = mesh.offset_provider["V2E"].asnumpy() + v2e_arr = mesh.offset_provider[cases_utils.V2EDim.tag].asnumpy() return cases.Case( ( None @@ -85,7 +93,7 @@ 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, "TaggedOff", TaggedOffDim) + return _case(exec_alloc_descriptor, TaggedOffDim.tag, TaggedOffDim) @pytest.fixture @@ -110,7 +118,7 @@ def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: cases.verify_with_default_data( case_tag_vs_variable_name, foo, - lambda a: a[_neighbor_table(case_tag_vs_variable_name, "TaggedOff")[:, 1]], + lambda a: a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)[:, 1]], ) @@ -122,7 +130,7 @@ def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: cases.verify_with_default_data( case_tag_vs_variable_name, foo, - lambda a: np.sum(a[_neighbor_table(case_tag_vs_variable_name, "TaggedOff")], axis=1), + lambda a: np.sum(a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)], axis=1), ) diff --git a/tests/next_tests/toy_connectivity.py b/tests/next_tests/toy_connectivity.py index 154b666c5d..c368b04245 100644 --- a/tests/next_tests/toy_connectivity.py +++ b/tests/next_tests/toy_connectivity.py @@ -12,18 +12,31 @@ from gt4py.next.iterator import builtins, ir as itir -Vertex = gtx.Dimension("Vertex") -Edge = gtx.Dimension("Edge") -Cell = gtx.Dimension("Cell") -V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) -E2VDim = gtx.Dimension("E2V", kind=gtx.DimensionKind.LOCAL) -C2EDim = gtx.Dimension("C2E", kind=gtx.DimensionKind.LOCAL) -V2VDim = gtx.Dimension("V2V", kind=gtx.DimensionKind.LOCAL) - -V2E = gtx.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset("E2V", source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset("C2E", source=Edge, target=(Cell, C2EDim)) -V2V = gtx.FieldOffset("V2V", source=Vertex, target=(Vertex, V2VDim)) +class Vertex(gtx.DimensionIndex): ... + + +class Edge(gtx.DimensionIndex): ... + + +class Cell(gtx.DimensionIndex): ... + + +class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +class V2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) +E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) +C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +V2V = gtx.FieldOffset(V2VDim.tag, source=Vertex, target=(Vertex, V2VDim)) # 3x3 periodic edges cells # 0 - 1 - 2 - 0 1 2 diff --git a/tests/next_tests/unit_tests/conftest.py b/tests/next_tests/unit_tests/conftest.py index 879094f123..edbd50c30a 100644 --- a/tests/next_tests/unit_tests/conftest.py +++ b/tests/next_tests/unit_tests/conftest.py @@ -23,6 +23,12 @@ import next_tests +class dummy_origin(gtx.DimensionIndex): ... + + +class dummy_neighbor(gtx.DimensionIndex): ... + + ProgramProcessor: TypeAlias = backend.Backend | program_formatter.ProgramFormatter @@ -95,8 +101,8 @@ def run_processor( class DummyConnectivity(common.Connectivity): max_neighbors: int has_skip_values: int - source_dim: gtx.Dimension = gtx.Dimension("dummy_origin") - codomain: gtx.Dimension = gtx.Dimension("dummy_neighbor") + source_dim: gtx.Dimension = dummy_origin + codomain: gtx.Dimension = dummy_neighbor def nd_array_implementation_params(): diff --git a/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py b/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py index a3e2c18e34..f5a3212889 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py @@ -11,7 +11,7 @@ import gt4py.next as gtx -IDim = gtx.Dimension("IDim") +class IDim(gtx.DimensionIndex): ... @gtx.field_operator diff --git a/tests/next_tests/unit_tests/embedded_tests/test_common.py b/tests/next_tests/unit_tests/embedded_tests/test_common.py index 57f9ef5648..2cf2cf3b47 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_common.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_common.py @@ -11,7 +11,7 @@ import pytest from gt4py.next import common -from gt4py.next.common import UnitRange, NamedIndex, NamedRange +from gt4py.next.common import UnitRange, DimensionIndex, NamedRange from gt4py.next.embedded import exceptions as embedded_exceptions from gt4py.next.embedded.common import ( _slice_range, @@ -37,9 +37,13 @@ def test_slice_range(rng, slce, expected): assert result == expected -I = common.Dimension("I") -J = common.Dimension("J") -K = common.Dimension("K") +class I(common.DimensionIndex): ... + + +class J(common.DimensionIndex): ... + + +class K(common.DimensionIndex): ... @pytest.mark.parametrize( @@ -47,11 +51,11 @@ def test_slice_range(rng, slce, expected): [ ([(I, (2, 5))], 1, []), ([(I, (2, 5))], slice(1, 2), [(I, (3, 4))]), - ([(I, (2, 5))], NamedIndex(I, 2), []), + ([(I, (2, 5))], I(2), []), ([(I, (2, 5))], NamedRange(I, UnitRange(2, 3)), [(I, (2, 3))]), ([(I, (-2, 3))], 1, []), ([(I, (-2, 3))], slice(1, 2), [(I, (-1, 0))]), - ([(I, (-2, 3))], NamedIndex(I, 1), []), + ([(I, (-2, 3))], I(1), []), ([(I, (-2, 3))], NamedRange(I, UnitRange(2, 3)), [(I, (2, 3))]), ([(I, (-2, 3))], -5, []), ([(I, (-2, 3))], -6, IndexError), @@ -61,9 +65,9 @@ def test_slice_range(rng, slce, expected): ([(I, (-2, 3))], 5, IndexError), ([(I, (-2, 3))], slice(4, 5), [(I, (2, 3))]), ([(I, (-2, 3))], slice(5, 6), IndexError), - ([(I, (-2, 3))], NamedIndex(I, -3), IndexError), + ([(I, (-2, 3))], I(-3), IndexError), ([(I, (-2, 3))], NamedRange(I, UnitRange(-3, -2)), IndexError), - ([(I, (-2, 3))], NamedIndex(I, 3), IndexError), + ([(I, (-2, 3))], I(3), IndexError), ([(I, (-2, 3))], NamedRange(I, UnitRange(3, 4)), IndexError), ([(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], 2, [(J, (3, 6)), (K, (4, 7))]), ( @@ -71,13 +75,13 @@ def test_slice_range(rng, slce, expected): slice(2, 3), [(I, (4, 5)), (J, (3, 6)), (K, (4, 7))], ), - ([(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], NamedIndex(I, 2), [(J, (3, 6)), (K, (4, 7))]), + ([(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], I(2), [(J, (3, 6)), (K, (4, 7))]), ( [(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], NamedRange(I, UnitRange(2, 3)), [(I, (2, 3)), (J, (3, 6)), (K, (4, 7))], ), - ([(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], NamedIndex(J, 3), [(I, (2, 5)), (K, (4, 7))]), + ([(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], J(3), [(I, (2, 5)), (K, (4, 7))]), ( [(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], NamedRange(J, UnitRange(4, 5)), @@ -85,12 +89,12 @@ def test_slice_range(rng, slce, expected): ), ( [(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], - (NamedIndex(J, 3), NamedIndex(I, 2)), + (J(3), I(2)), [(K, (4, 7))], ), ( [(I, (2, 5)), (J, (3, 6)), (K, (4, 7))], - (NamedRange(J, UnitRange(4, 5)), NamedIndex(I, 2)), + (NamedRange(J, UnitRange(4, 5)), I(2)), [(J, (4, 5)), (K, (4, 7))], ), ( @@ -131,7 +135,7 @@ def test_iterate_domain(): ref = [] for i in domain[I].unit_range: for j in domain[J].unit_range: - ref.append(((I, i), (J, j))) + ref.append((I(i), J(j))) testee = list(iterate_domain(domain)) diff --git a/tests/next_tests/unit_tests/embedded_tests/test_context.py b/tests/next_tests/unit_tests/embedded_tests/test_context.py index 8ceeb82174..165ec82b2e 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_context.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_context.py @@ -13,6 +13,12 @@ from gt4py.next.errors import exceptions +class IDim(common.DimensionIndex): ... + + +class NewDim(common.DimensionIndex): ... + + def test_getters(): DEFAULT = object() assert ctx.get_closure_column_range(DEFAULT) is DEFAULT @@ -30,7 +36,7 @@ def test_update_with_both_parameters(): assert ctx.get_closure_column_range(DEFAULT) is DEFAULT assert ctx.get_offset_provider(DEFAULT) is DEFAULT - initial_column_range = common.NamedRange(common.Dimension("IDim"), common.UnitRange(0, 4)) + initial_column_range = common.NamedRange(IDim, common.UnitRange(0, 4)) initial_offset_provider = {} with ctx.update( @@ -39,7 +45,7 @@ def test_update_with_both_parameters(): assert ctx.get_closure_column_range() is initial_column_range assert ctx.get_offset_provider() is initial_offset_provider - test_column_range = common.NamedRange(common.Dimension("NewDim"), common.UnitRange(-1, 1)) + test_column_range = common.NamedRange(NewDim, common.UnitRange(-1, 1)) test_offset_provider = {"I": "NewDim"} with ctx.update( @@ -60,7 +66,7 @@ def test_update_with_no_parameters(): assert ctx.get_closure_column_range(DEFAULT) is DEFAULT assert ctx.get_offset_provider(DEFAULT) is DEFAULT - initial_column_range = common.NamedRange(common.Dimension("IDim"), common.UnitRange(0, 4)) + initial_column_range = common.NamedRange(IDim, common.UnitRange(0, 4)) initial_offset_provider = {} with ctx.update( @@ -85,7 +91,7 @@ def test_update_with_exception(): assert ctx.get_closure_column_range(DEFAULT) is DEFAULT assert ctx.get_offset_provider(DEFAULT) is DEFAULT - initial_column_range = common.NamedRange(common.Dimension("IDim"), common.UnitRange(0, 4)) + initial_column_range = common.NamedRange(IDim, common.UnitRange(0, 4)) initial_offset_provider = {} with pytest.raises(RuntimeError, match="Outer exception"): @@ -95,9 +101,7 @@ def test_update_with_exception(): assert ctx.get_closure_column_range() is initial_column_range assert ctx.get_offset_provider() is initial_offset_provider - test_column_range = common.NamedRange( - common.Dimension("NewDim"), common.UnitRange(-1, 1) - ) + test_column_range = common.NamedRange(NewDim, common.UnitRange(-1, 1)) test_offset_provider = {"I": "NewDim"} with pytest.raises(RuntimeError, match="Inner exception"): 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 3546ceaaaa..bef6a29c8a 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 @@ -18,10 +18,11 @@ from gt4py.next import common, constructors from gt4py.next.common import ( Dimension, + DimensionIndex, DimensionKind, Domain, Field, - NamedIndex, + DimensionIndex, NamedRange, UnitRange, ) @@ -33,9 +34,82 @@ from next_tests.integration_tests.feature_tests.math_builtin_test_data import math_builtin_test_data -D0 = Dimension("D0") -D1 = Dimension("D1") -D2 = Dimension("D2") +class I(DimensionIndex): ... + + +class J(DimensionIndex): ... + + +class I_half(DimensionIndex): ... + + +class V(DimensionIndex): ... + + +class E(DimensionIndex): ... + + +class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class V2V(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class C(DimensionIndex): ... + + +class K(DimensionIndex): ... + + +class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class A(DimensionIndex): ... + + +class B(DimensionIndex): ... + + +class X(DimensionIndex): ... + + +class Y(DimensionIndex): ... + + +class L(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class S(DimensionIndex): ... + + +class T(DimensionIndex): ... + + +class U(DimensionIndex): ... + + +class M(DimensionIndex): ... + + +class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + + +class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + + +class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class C2V(DimensionIndex): ... + + +class D0(DimensionIndex): ... + + +class D1(DimensionIndex): ... + + +class D2(DimensionIndex): ... @pytest.fixture( @@ -67,9 +141,24 @@ def unary_logical_op(request): yield request.param +def _default_dim(i: int) -> Dimension: + """ + Return the `i`-th default dimension, `D`, creating it on first use. + + The dimensions must be *stable across calls*: two domains built by separate calls are + compared by the tests, and under nominal identity (ADR 0028) a fresh class per call would + be a different dimension. Each is bound as a module attribute, so it is also importable + by its qualified name like any other dimension. + """ + name = f"D{i}" + if name not in globals(): + globals()[name] = type(name, (common.DimensionIndex,), {"__module__": __name__}) + return globals()[name] + + def _make_default_domain(shape: tuple[int, ...]) -> Domain: return common.Domain( - dims=tuple(Dimension(f"D{i}") for i in range(len(shape))), + dims=tuple(_default_dim(i) for i in range(len(shape))), ranges=tuple(UnitRange(0, s) for s in shape), ) @@ -348,8 +437,6 @@ def fma(a: common.Field, b: common.Field, c: common.Field, /) -> common.Field: def test_domain_premap(): # Translation case - I = Dimension("I") - J = Dimension("J") N = 10 data_field = common._field( @@ -373,7 +460,6 @@ def test_domain_premap(): assert np.all(result.ndarray == expected.ndarray) # Relocation case - I_half = Dimension("I_half") conn = common.CartesianConnectivity.for_relocation(I, I_half) @@ -394,8 +480,6 @@ def test_domain_premap(): def test_reshuffling_premap(): - I = Dimension("I") - J = Dimension("J") ij_field = common._field( np.asarray([[0.0, 1.0, 2.0], [3.0, 4.0, 5.0], [6.0, 7.0, 8.0]]), @@ -422,8 +506,6 @@ def test_reshuffling_premap(): def test_remapping_premap(): - V = Dimension("V") - E = Dimension("E") V_START, V_STOP = 2, 7 E_START, E_STOP = 0, 10 @@ -448,9 +530,6 @@ def test_remapping_premap(): def test_remapping_premap_multineighbor(): - V = Dimension("V") - E = Dimension("E") - E2V = Dimension("E2V", kind=DimensionKind.LOCAL) V_START, V_STOP = 2, 7 v_field = common._field( @@ -475,7 +554,6 @@ def test_remapping_premap_multineighbor(): def test_remapping_premap_same_dim(): - V = Dimension("V") V_START, V_STOP = 2, 7 v_field = common._field( @@ -500,8 +578,6 @@ def test_remapping_premap_same_dim(): def test_remapping_premap_same_dim_multineighbor(): # Regression for #1583: source dim == target dim with a LOCAL neighbor dim (C2E2CO shape). - V = Dimension("V") - V2V = Dimension("V2V", kind=DimensionKind.LOCAL) V_START, V_STOP = 2, 7 v_field = common._field( @@ -535,9 +611,6 @@ def test_remapping_premap_same_dim_multineighbor(): def test_premap_same_dim_multineighbor_with_extra_dim(): # The actual #1583 bug shape: C2E2CO on a field with an extra (kept) dimension. - C = Dimension("C") - K = Dimension("K") - C2E2CO = Dimension("C2E2CO", kind=DimensionKind.LOCAL) NC, NK, NN = 5, 3, 4 c_field = common._field( @@ -561,10 +634,6 @@ def test_premap_same_dim_multineighbor_with_extra_dim(): def test_gather_premap_multiple_connectivities(): # Multiple gather connectivities introducing new dims in a single premap (was unsupported). - A = Dimension("A") - B = Dimension("B") - X = Dimension("X") - Y = Dimension("Y") NA, NB = 3, 4 f = common._field( @@ -589,8 +658,6 @@ def test_gather_premap_multiple_connectivities(): def test_gather_premap_reshuffle_multiple_connectivities(): # Simultaneous multi-axis reshuffle (the multi-axis `as_offset` shape): codomains stay in their # own domains, so out[i, j] = f[ci[i, j], cj[i, j]] (not a sequential composition). - I = Dimension("I") - J = Dimension("J") NI, NJ = 3, 4 f = common._field( @@ -611,9 +678,6 @@ def test_gather_premap_reshuffle_multiple_connectivities(): def test_gather_premap_two_connectivities_same_new_dim(): # Two connectivities introducing the *same* new dim -> diagonal gather out[x] = f[ca[x], cb[x]]. - A = Dimension("A") - B = Dimension("B") - X = Dimension("X") NA, NB = 3, 4 f = common._field( @@ -637,9 +701,6 @@ def test_gather_premap_two_connectivities_same_new_dim(): def test_gather_premap_reads_non_codomain_field_dim(): # A connectivity that reads a field dim (B) other than its codomain (A): B is shared/narrowed. - A = Dimension("A") - B = Dimension("B") - L = Dimension("L", kind=DimensionKind.LOCAL) NA, NB = 3, 4 f = common._field( @@ -662,11 +723,6 @@ def test_gather_premap_reads_non_codomain_field_dim(): def test_gather_premap_shared_domain_dim(): # Two connectivities sharing a non-codomain domain dim (S): S is intersected and kept once. - A = Dimension("A") - B = Dimension("B") - S = Dimension("S") - T = Dimension("T") - U = Dimension("U") NA, NB = 3, 4 f = common._field( @@ -693,9 +749,6 @@ def test_gather_premap_shared_domain_dim(): def test_gather_premap_mix_introducing_and_preserving(): # One connectivity introduces a dim (X), another reshuffles a dim in place (B): allowed together. - A = Dimension("A") - B = Dimension("B") - X = Dimension("X") NA, NB = 3, 4 f = common._field( @@ -720,10 +773,6 @@ def test_gather_premap_mix_introducing_and_preserving(): def test_premap_chained_connectivities_raises(): # One connectivity reads a dim (B) that another remaps -> unsupported chained composition. - A = Dimension("A") - B = Dimension("B") - L = Dimension("L", kind=DimensionKind.LOCAL) - M = Dimension("M") f = common._field( np.arange(12).reshape(3, 4).astype(float), @@ -746,8 +795,6 @@ def test_premap_chained_connectivities_raises(): def test_premap_non_contiguous_inverse_image_raises(): # A connectivity whose in-range indices are not a contiguous block cannot yield a contiguous domain. - V = Dimension("V") - E = Dimension("E") f = common._field( np.arange(5).astype(float), domain=common.Domain(dims=(V,), ranges=(UnitRange(0, 5),)) @@ -763,8 +810,6 @@ def test_premap_non_contiguous_inverse_image_raises(): def test_premap_disjoint_inverse_image_raises(): - V = Dimension("V") - E = Dimension("E") f = common._field( np.arange(5).astype(float), domain=common.Domain(dims=(V,), ranges=(UnitRange(0, 5),)) @@ -781,7 +826,6 @@ def test_premap_disjoint_inverse_image_raises(): def test_as_offset_1d(): # Dynamic per-point shift along I: out[i] == f[i + off[i]], full domain when all shifts in-bounds. - I = Dimension("I") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( @@ -798,7 +842,6 @@ def test_as_offset_1d(): def test_as_offset_narrow_offset_dtype_no_wrap(): # An int8 offset field over a domain larger than 128 must not wrap into the index table. - I = Dimension("I") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) N = 200 @@ -816,8 +859,6 @@ def test_as_offset_narrow_offset_dtype_no_wrap(): def test_as_offset_2d_shift_one_keep_other(): # Shift along I by a per-(i, j) offset, leave J: out[i, j] == f[i + off[i, j], j]. - I = Dimension("I") - J = Dimension("J") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) NI, NJ = 4, 3 @@ -836,7 +877,6 @@ def test_as_offset_2d_shift_one_keep_other(): def test_as_offset_boundary_narrows_domain(): # A uniform out-of-bounds shift narrows the result to the contiguous in-range sub-domain. - I = Dimension("I") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( @@ -854,7 +894,6 @@ def test_as_offset_boundary_narrows_domain(): def test_as_offset_scattered_oob_raises(): # An out-of-bounds shift in the interior cannot yield a contiguous domain. - I = Dimension("I") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( @@ -871,8 +910,6 @@ def test_as_offset_scattered_oob_raises(): def test_as_offset_introduces_dimension(): # `off` carries a dim the field lacks: the result gains it, out[i, j] == f[i + off[i, j]]. - I = Dimension("I") - J = Dimension("J") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( @@ -891,7 +928,6 @@ def test_as_offset_introduces_dimension(): def test_as_offset_nonzero_origin(): # Field and offset over a domain that does not start at 0: indices must be shifted by the domain start. - I = Dimension("I") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) dom = common.Domain(dims=(I,), ranges=(UnitRange(2, 12),)) @@ -907,8 +943,6 @@ def test_as_offset_nonzero_origin(): def test_as_offset_2d_shift_second_axis(): # Shift along J (the non-leading axis) by a per-(i, j) offset, leave I: out[i, j] == f[i, j + off[i, j]]. - I = Dimension("I") - J = Dimension("J") Joff = fbuiltins.FieldOffset("Joff", source=J, target=(J,)) NI, NJ = 3, 4 @@ -927,11 +961,6 @@ def test_as_offset_2d_shift_second_axis(): def test_as_offset_non_cartesian_offset_raises(): # `as_offset` only supports Cartesian (self-shift) offsets: single target equal to source. - I = Dimension("I") - J = Dimension("J") - Vertex = Dimension("Vertex", kind=DimensionKind.HORIZONTAL) - Edge = Dimension("Edge", kind=DimensionKind.HORIZONTAL) - V2EDim = Dimension("V2EDim", kind=DimensionKind.LOCAL) off_I = common._field( np.zeros(3, dtype=int), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 3),)) @@ -941,7 +970,7 @@ def test_as_offset_non_cartesian_offset_raises(): ) # 2-element target (neighbor offset) - V2E = fbuiltins.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) + V2E = fbuiltins.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) with pytest.raises(ValueError, match="Cartesian"): as_offset(V2E, off_V) @@ -1016,7 +1045,7 @@ def test_get_slices_with_named_index(): field_domain = common.Domain( dims=(D0, D1, D2), ranges=(UnitRange(0, 10), UnitRange(0, 10), UnitRange(0, 10)) ) - named_index = (NamedRange(D0, UnitRange(0, 10)), (D1, 2), (D2, 3)) + named_index = (NamedRange(D0, UnitRange(0, 10)), D1(2), D2(3)) slices = _get_slices_from_domain_slice(field_domain, named_index) assert slices == (slice(0, 10, None), 2, 3) @@ -1025,7 +1054,7 @@ def test_get_slices_invalid_type(): field_domain = common.Domain( dims=(D0, D1, D2), ranges=(UnitRange(0, 10), UnitRange(0, 10), UnitRange(0, 10)) ) - new_domain = ((D0, "1"),) + new_domain = (D0("1"),) with pytest.raises(ValueError): _get_slices_from_domain_slice(field_domain, new_domain) @@ -1044,11 +1073,11 @@ def test_get_slices_invalid_type(): (2, 10, 8), ), (common.Domain(dims=(D0,), ranges=(UnitRange(7, 9),)), (D0, D1, D2), (2, 10, 15)), - ((NamedIndex(D0, 8),), (D1, D2), (10, 15)), - ((NamedIndex(D1, 9),), (D0, D2), (5, 15)), - ((NamedIndex(D2, 11),), (D0, D1), (5, 10)), - ((NamedIndex(D0, 8), NamedRange(D1, UnitRange(8, 10))), (D1, D2), (2, 15)), - (NamedIndex(D0, 5), (D1, D2), (10, 15)), + ((D0(8),), (D1, D2), (10, 15)), + ((D1(9),), (D0, D2), (5, 15)), + ((D2(11),), (D0, D1), (5, 10)), + ((D0(8), NamedRange(D1, UnitRange(8, 10))), (D1, D2), (2, 15)), + (D0(5), (D1, D2), (10, 15)), (NamedRange(D0, UnitRange(5, 7)), (D0, D1, D2), (2, 10, 15)), ], ) @@ -1085,7 +1114,7 @@ def test_absolute_indexing_dim_sliced_single_slice(): ) field = common._field(np.ones((5, 10, 15)), domain=domain) indexed_field_1 = field[D2(11)] - indexed_field_2 = field[NamedIndex(D2, 11)] + indexed_field_2 = field[D2(11)] assert isinstance(indexed_field_1, common.Field) assert are_equal_fields(indexed_field_1, indexed_field_2) @@ -1114,7 +1143,7 @@ def test_absolute_indexing_value_return(): domain = common.Domain(dims=(D0, D1), ranges=(UnitRange(10, 20), UnitRange(5, 15))) field = common._field(np.reshape(np.arange(100, dtype=np.int32), (10, 10)), domain=domain) - named_index = (NamedIndex(D0, 12), NamedIndex(D1, 6)) + named_index = (D0(12), D1(6)) assert isinstance(field, common.Field) value = field[named_index] @@ -1316,9 +1345,6 @@ def test_nd_array_field_pickle_roundtrip(): def test_nd_array_connectivity_field_buffer_info(nd_array_implementation): import dataclasses - V = Dimension("V") - E = Dimension("E") - V_START, V_STOP = 2, 7 E_START, E_STOP = 0, 10 @@ -1334,8 +1360,6 @@ def test_nd_array_connectivity_field_buffer_info(nd_array_implementation): def test_nd_array_connectivity_field_getstate_excludes_runtime_caches(): - V = Dimension("V") - E = Dimension("E") e2v_conn = common._connectivity( np.asarray([2, 3, 4, 5]), @@ -1357,8 +1381,6 @@ def test_nd_array_connectivity_field_getstate_excludes_runtime_caches(): def test_nd_array_connectivity_field_setstate_restores_state_without_caches(): - V = Dimension("V") - E = Dimension("E") original = common._connectivity( np.asarray([2, 3, 4, 5]), @@ -1382,8 +1404,6 @@ def test_nd_array_connectivity_field_setstate_restores_state_without_caches(): def test_connectivity_field_inverse_image(): - V = Dimension("V") - E = Dimension("E") V_START, V_STOP = 2, 7 E_START, E_STOP = 0, 10 @@ -1396,9 +1416,6 @@ def test_connectivity_field_inverse_image(): def test_connectivity_field_inverse_image_2d_domain(): - V = Dimension("V") - C = Dimension("C") - C2V = Dimension("C2V") V_START, V_STOP = 0, 3 C_START, C_STOP = 0, 3 @@ -1454,8 +1471,6 @@ def test_connectivity_field_inverse_image_2d_domain(): def test_connectivity_field_inverse_image_non_contiguous(): - V = Dimension("V") - E = Dimension("E") V_START, V_STOP = 2, 7 E_START, E_STOP = 0, 10 @@ -1477,9 +1492,6 @@ def test_connectivity_field_inverse_image_non_contiguous(): def test_connectivity_field_inverse_image_2d_domain_skip_values(): - V = Dimension("V") - C = Dimension("C") - C2V = Dimension("C2V") V_START, V_STOP = 0, 3 C_START, C_STOP = 0, 4 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 82abd8a7a7..25284281ef 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 @@ -12,11 +12,20 @@ from gt4py.next.ffront.transform_utils import _deduce_grid_type -Dim = gtx.Dimension("Dim") -LocalDim = gtx.Dimension("LocalDim", kind=gtx.DimensionKind.LOCAL) +class HDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... + + +class VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class Dim(gtx.DimensionIndex): ... + + +class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + CartesianOffset = gtx.FieldOffset("CartesianOffset", source=Dim, target=(Dim,)) -UnstructuredOffset = gtx.FieldOffset("UnstructuredOffset", source=Dim, target=(Dim, LocalDim)) +UnstructuredOffset = gtx.FieldOffset(LocalDim.tag, source=Dim, target=(Dim, LocalDim)) def test_domain_deduction_cartesian(): @@ -28,8 +37,6 @@ def test_domain_deduction_unstructured(): assert _deduce_grid_type(None, {UnstructuredOffset}) == gtx.GridType.UNSTRUCTURED assert _deduce_grid_type(None, {LocalDim}) == gtx.GridType.UNSTRUCTURED # source and target share `.value` but differ in `.kind` -> not Cartesian - HDim = gtx.Dimension("X", kind=gtx.DimensionKind.HORIZONTAL) - VDim = gtx.Dimension("X", kind=gtx.DimensionKind.VERTICAL) CrossKindOffset = gtx.FieldOffset("CrossKind", source=HDim, target=(VDim,)) assert _deduce_grid_type(None, {CrossKindOffset}) == gtx.GridType.UNSTRUCTURED # LOCAL self-loop is unstructured diff --git a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py index 7e2d4fed6e..337c7ce67f 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py @@ -27,7 +27,9 @@ from gt4py.next.ffront.func_to_foast import FieldOperatorParser -IDim = gtx.Dimension("IDim") +class IDim(gtx.DimensionIndex): ... + + IOff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) # A PEP 695 alias whose value raises when it is evaluated, standing in for the diff --git a/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py b/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py index f9d852e730..9018742fe0 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py @@ -20,7 +20,8 @@ # values inside the domain of every unary math builtin (0.5 is invalid for `arccosh`) _SAFE_INPUT = {"arccosh": 2.0} -IDim = common.Dimension("IDim") + +class IDim(common.DimensionIndex): ... @dataclasses.dataclass 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 eb1927a776..f696daf1b4 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 @@ -43,18 +43,33 @@ from gt4py.next.iterator import ir as itir -Edge = gtx.Dimension("Edge") -Vertex = gtx.Dimension("Vertex") -V2EDim = gtx.Dimension("V2E", gtx.DimensionKind.LOCAL) -V2E = gtx.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) +class Edge(gtx.DimensionIndex): ... + + +class Vertex(gtx.DimensionIndex): ... + + +class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) + + +class TDim(gtx.DimensionIndex): ... + -TDim = gtx.Dimension("TDim") TOff = gtx.FieldOffset("TDim", source=TDim, target=(TDim,)) + + #: An offset whose tag differs from the name of the Python variable it is bound to, and #: from the name of its local dimension. Lowering must emit the *tag*. -RenamedV2EDim = gtx.Dimension("RenamedLocal", gtx.DimensionKind.LOCAL) +class RenamedV2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) -UDim = gtx.Dimension("UDim") + + +class UDim(gtx.DimensionIndex): ... def test_return(): @@ -781,7 +796,7 @@ def foo(edge_f: gtx.Field[gtx.Dims[Edge], float64]): parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) - reference = im.as_fieldop_neighbors("V2E", "edge_f") + reference = im.as_fieldop_neighbors(V2EDim.tag, "edge_f") assert lowered.expr == reference @@ -826,7 +841,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): "plus", im.literal(value="0", type_="float64"), ) - )(im.as_fieldop_neighbors("V2E", "edge_f")) + )(im.as_fieldop_neighbors(V2EDim.tag, "edge_f")) assert lowered.expr == reference @@ -843,7 +858,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): "maximum", im.literal(value=str(np.finfo(np.float64).min), type_="float64"), ) - )(im.as_fieldop_neighbors("V2E", "edge_f")) + )(im.as_fieldop_neighbors(V2EDim.tag, "edge_f")) assert lowered.expr == reference @@ -860,7 +875,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): "minimum", im.literal(value=str(np.finfo(np.float64).max), type_="float64"), ) - )(im.as_fieldop_neighbors("V2E", "edge_f")) + )(im.as_fieldop_neighbors(V2EDim.tag, "edge_f")) assert lowered.expr == reference @@ -880,7 +895,7 @@ def foo(e1: gtx.Field[[Edge], float64], e2: gtx.Field[[Vertex, V2EDim], float64] reference = im.let( ssa.unique_name("e1_nbh", 0), - im.as_fieldop_neighbors("V2E", "e1"), + im.as_fieldop_neighbors(V2EDim.tag, "e1"), )( im.op_as_fieldop( im.reduce( @@ -970,7 +985,7 @@ def foo(inp: gtx.Field[[TDim], float64]): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( im.ref("inp"), - im.make_tuple(*(itir.AxisLiteral(value=dim.value, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), ) @@ -984,7 +999,7 @@ def foo(): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( 1, - im.make_tuple(*(itir.AxisLiteral(value=dim.value, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), ) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py index d72efaaa22..7c3786c28f 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py @@ -54,6 +54,12 @@ from gt4py.next.type_system import type_specifications as ts +class ADim(gtx.DimensionIndex): ... + + +class BDim(gtx.DimensionIndex): ... + + DEREF = itir.SymRef(id=itb.deref.fun.__name__) PLUS = itir.SymRef(id=itb.plus.fun.__name__) MINUS = itir.SymRef(id=itb.minus.fun.__name__) @@ -71,7 +77,9 @@ XOR = itir.SymRef(id=itb.xor_.fun.__name__) LIFT = itir.SymRef(id=itb.lift.fun.__name__) -TDim = gtx.Dimension("TDim") # Meaningless dimension, used for tests. + +class TDim(gtx.DimensionIndex): ... + # PEP 695 type alias, used to check that aliases are accepted as DSL annotations. type TFloatFieldAlias = gtx.Field[gtx.Dims[TDim], float64] @@ -426,8 +434,6 @@ def operator_with_refs(inp2: gtx.Field[[TDim], "float32"]): def test_wrong_return_type_annotation(): - ADim = gtx.Dimension("ADim") - BDim = gtx.Dimension("BDim") def wrong_return_type_annotation(a: gtx.Field[[ADim], float64]) -> gtx.Field[[BDim], float64]: return a @@ -449,7 +455,6 @@ def empty_dims() -> gtx.Field[[], float]: def test_zero_dims_ternary(): - ADim = gtx.Dimension("ADim") def zero_dims_ternary( cond: gtx.Field[[], float64], a: gtx.Field[[ADim], float64], b: gtx.Field[[ADim], float64] diff --git a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py index 62a8f1b755..330436646b 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py @@ -19,7 +19,8 @@ # NOTE: These tests are sensitive to filename and the line number of the marked statement -TDim = gtx.Dimension("TDim") # Meaningless dimension, used for tests. + +class TDim(gtx.DimensionIndex): ... def test_invalid_syntax_error_empty_return(): diff --git a/tests/next_tests/unit_tests/ffront_tests/test_past_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_past_to_gtir.py index 9db3567a2f..41943b5dbb 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_past_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_past_to_gtir.py @@ -78,9 +78,9 @@ def test_copy_lowering(copy_program_def, gtir_identity_fundef): itir.FunCall, fun=P(itir.SymRef, id=eve.SymbolRef("named_range")), args=[ - P(itir.AxisLiteral, value="IDim"), - get_domain_range_pattern("out", "IDim", 0), - get_domain_range_pattern("out", "IDim", 1), + P(itir.AxisLiteral, value=IDim.tag), + get_domain_range_pattern("out", IDim.tag, 0), + get_domain_range_pattern("out", IDim.tag, 1), ], ) ], @@ -122,12 +122,12 @@ def test_copy_restrict_lowering(copy_restrict_program_def, gtir_identity_fundef) itir.FunCall, fun=P(itir.SymRef, id=eve.SymbolRef("named_range")), args=[ - P(itir.AxisLiteral, value="IDim"), + P(itir.AxisLiteral, value=IDim.tag), P( itir.FunCall, fun=P(itir.SymRef, id=eve.SymbolRef("plus")), args=[ - get_domain_range_pattern("out", "IDim", 0), + get_domain_range_pattern("out", IDim.tag, 0), P( itir.Literal, value="1", @@ -143,7 +143,7 @@ def test_copy_restrict_lowering(copy_restrict_program_def, gtir_identity_fundef) itir.FunCall, fun=P(itir.SymRef, id=eve.SymbolRef("plus")), args=[ - get_domain_range_pattern("out", "IDim", 0), + get_domain_range_pattern("out", IDim.tag, 0), P( itir.Literal, value="2", 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 5735b4c25b..290d2914fc 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,15 +19,21 @@ """ import gt4py.next as gtx -from gt4py.next import Dims, Dimension, float64, neighbor_sum +from gt4py.next import Dims, Dimension, DimensionIndex, float64, neighbor_sum from gt4py.next.ffront import source_utils from gt4py.next.ffront.source_utils import get_closure_vars_from_function -Cell = Dimension("Cell") -Edge = Dimension("Edge") -C2EDim = Dimension("C2E", kind=gtx.DimensionKind.LOCAL) -C2E = gtx.FieldOffset("C2E", source=Edge, target=(Cell, C2EDim)) +class Cell(DimensionIndex): ... + + +class Edge(DimensionIndex): ... + + +class C2EDim(DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) CField = gtx.Field[Dims[Cell], float64] EField = gtx.Field[Dims[Edge], float64] diff --git a/tests/next_tests/unit_tests/ffront_tests/test_stages.py b/tests/next_tests/unit_tests/ffront_tests/test_stages.py index 3c14182655..3509797241 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_stages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_stages.py @@ -15,7 +15,7 @@ from gt4py.next.type_system import type_specifications as ts -IDim = gtx.Dimension("I") +class IDim(gtx.DimensionIndex): ... def _field_type(kind: ts.ScalarKind) -> ts.FieldType: 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 b6ae674312..88a8c640d8 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 @@ -16,6 +16,7 @@ import gt4py.next.ffront.type_specifications from gt4py.next import ( Dimension, + DimensionIndex, DimensionKind, Field, FieldOffset, @@ -37,9 +38,50 @@ from next_tests.artifacts import custom_named_collections as cnc +# NOTE: the named collections in the artifact are annotated with *its* `TDim`, and the expected +# types below are compared against them. Under nominal identity (ADR 0028) a redeclared `TDim` +# here would be a different dimension, where the old `Dimension("TDim")` compared equal. +TDim = cnc.TDim + + +class X(DimensionIndex): ... + + +class Y(DimensionIndex): ... + + +class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... + + +class ADim(DimensionIndex): ... + + +class BDim(DimensionIndex): ... + + +class CDim(DimensionIndex): ... + + +class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + + +class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + + +class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class IDim(DimensionIndex): ... + + +class JDim(DimensionIndex): ... + + # Meaningless dimensions, used for tests. -TDim = Dimension("TDim") -SDim = Dimension("SDim") +class SDim(DimensionIndex): ... def test_unpack_assign(): @@ -93,8 +135,6 @@ def add_bools(a: Field[[TDim], bool], b: Field[[TDim], bool]): def test_binop_nonmatching_dims(): """Dimension promotion is applied before Binary operations, i.e., they can also work on two fields that don't have the same dimensions.""" - X = Dimension("X") - Y = Dimension("Y") def nonmatching(a: Field[[X], float64], b: Field[[Y], float64]): return a + b @@ -246,10 +286,7 @@ def domain_comparison(a: Field[[TDim], float], b: Field[[TDim], float]): @pytest.fixture def premap_setup(): - X = Dimension("X") - Y = Dimension("Y") - Y2XDim = Dimension("Y2X", kind=DimensionKind.LOCAL) - Y2X = FieldOffset("Y2X", source=X, target=(Y, Y2XDim)) + Y2X = FieldOffset(Y2XDim.tag, source=X, target=(Y, Y2XDim)) return X, Y, Y2XDim, Y2X @@ -281,7 +318,6 @@ def premap_fo(bar: Field[[X], int64]) -> Field[[Y, Y2XDim], int64]: def test_premap_nbfield_with_vertical(premap_setup): X, Y, Y2XDim, Y2X = premap_setup - K = Dimension("K", kind=DimensionKind.VERTICAL) def premap_fo(bar: Field[[X, K], int64]) -> Field[[Y, Y2XDim, K], int64]: return bar(Y2X) @@ -331,9 +367,6 @@ def mismatched_lit() -> Field[[TDim], "float32"]: def test_broadcast_multi_dim(): - ADim = Dimension("ADim") - BDim = Dimension("BDim") - CDim = Dimension("CDim") def simple_broadcast(a: Field[[ADim], float64]): return broadcast(a, (ADim, BDim, CDim)) @@ -346,9 +379,6 @@ def simple_broadcast(a: Field[[ADim], float64]): def test_broadcast_disjoint(): - ADim = Dimension("ADim") - BDim = Dimension("BDim") - CDim = Dimension("CDim") def disjoint_broadcast(a: Field[[ADim], float64]): return broadcast(a, (BDim, CDim)) @@ -358,9 +388,7 @@ def disjoint_broadcast(a: Field[[ADim], float64]): def test_broadcast_badtype(): - ADim = Dimension("ADim") BDim = "BDim" - CDim = Dimension("CDim") def badtype_broadcast(a: Field[[ADim], float64]): return broadcast(a, (BDim, CDim)) @@ -372,8 +400,6 @@ def badtype_broadcast(a: Field[[ADim], float64]): def test_where_dim(): - ADim = Dimension("ADim") - BDim = Dimension("BDim") def simple_where(a: Field[[ADim], bool], b: Field[[ADim, BDim], float64]): return where(a, b, 9.0) @@ -386,7 +412,6 @@ def simple_where(a: Field[[ADim], bool], b: Field[[ADim, BDim], float64]): def test_where_broadcast_dim(): - ADim = Dimension("ADim") def simple_where(a: Field[[ADim], bool]): return where(a, 5.0, 9.0) @@ -399,7 +424,6 @@ def simple_where(a: Field[[ADim], bool]): def test_where_tuple_dim(): - ADim = Dimension("ADim") def tuple_where(a: Field[[ADim], bool], b: Field[[ADim], float64]): return where(a, ((5.0, 9.0), (b, 6.0)), ((8.0, b), (5.0, 9.0))) @@ -425,7 +449,6 @@ def tuple_where(a: Field[[ADim], bool], b: Field[[ADim], float64]): def test_where_bad_dim(): - ADim = Dimension("ADim") def bad_dim_where(a: Field[[ADim], bool], b: Field[[ADim], float64]): return where(a, ((5.0, 9.0), (b, 6.0)), b) @@ -438,8 +461,6 @@ def bad_dim_where(a: Field[[ADim], bool], b: Field[[ADim], float64]): def test_where_mixed_dims(): - ADim = Dimension("ADim") - BDim = Dimension("BDim") def tuple_where_mix_dims( a: Field[[ADim], bool], b: Field[[ADim], float64], c: Field[[ADim, BDim], float64] @@ -526,8 +547,6 @@ def return_undefined(): def test_as_offset_dim(): - ADim = Dimension("ADim") - BDim = Dimension("BDim") Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) def as_offset_dim(a: Field[[ADim, BDim], float], b: Field[[ADim], int]): @@ -538,8 +557,6 @@ def as_offset_dim(a: Field[[ADim, BDim], float], b: Field[[ADim], int]): def test_as_offset_dtype(): - ADim = Dimension("ADim") - BDim = Dimension("BDim") Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) def as_offset_dtype(a: Field[[ADim, BDim], float], b: Field[[BDim], float]): @@ -550,10 +567,7 @@ def as_offset_dtype(a: Field[[ADim, BDim], float], b: Field[[BDim], float]): def test_as_offset_non_cartesian(): - Vertex = Dimension("Vertex", kind=DimensionKind.HORIZONTAL) - Edge = Dimension("Edge", kind=DimensionKind.HORIZONTAL) - V2EDim = Dimension("V2EDim", kind=DimensionKind.LOCAL) - V2E = FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) + V2E = FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) def as_offset_neighbor(a: Field[[Edge], float], b: Field[[Edge], int]): return a(as_offset(V2E, b)) @@ -561,8 +575,6 @@ def as_offset_neighbor(a: Field[[Edge], float], b: Field[[Edge], int]): with pytest.raises(errors.DSLError, match="Cartesian"): _ = FieldOperatorParser.apply_to_function(as_offset_neighbor) - IDim = Dimension("IDim") - JDim = Dimension("JDim") IfromJ = FieldOffset("IfromJ", source=IDim, target=(JDim,)) def as_offset_cross_dim(a: Field[[IDim], float], b: Field[[IDim], int]): 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 4aac031406..7120882c75 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 @@ -14,15 +14,33 @@ from gt4py.next.iterator.ir_utils import domain_utils, ir_makers as im from gt4py.next import common, constructors -I = common.Dimension("I") + +class I(common.DimensionIndex): ... + + IHalf = common.flip_staggered(I) -J = common.Dimension("J") -K = common.Dimension("J", kind=common.DimensionKind.VERTICAL) -Vertex = common.Dimension("Vertex") -Edge = common.Dimension("Edge") -V2EDim = common.Dimension("V2E", kind=common.DimensionKind.LOCAL) -E2VDim = common.Dimension("E2V", kind=common.DimensionKind.LOCAL) -V2VDim = common.Dimension("V2V", kind=common.DimensionKind.LOCAL) + + +class J(common.DimensionIndex): ... + + +class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + + +class Vertex(common.DimensionIndex): ... + + +class Edge(common.DimensionIndex): ... + + +class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class V2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + a_range = domain_utils.SymbolicRange(0, 10) another_range = domain_utils.SymbolicRange(5, 15) @@ -103,7 +121,7 @@ def test_domain_union_all_empty(): def test_unstructured_translate_empty_range(): offset_provider = { - "V2E": constructors.as_connectivity( + V2EDim.tag: constructors.as_connectivity( domain={Vertex: (0, 4), V2EDim: 1}, codomain=Edge, data=np.asarray([0, 1, 2, 3], dtype=fbuiltins.IndexType).reshape((4, 1)), @@ -113,7 +131,7 @@ def test_unstructured_translate_empty_range(): im.domain(common.GridType.UNSTRUCTURED, {Vertex: (2, 2)}) # empty ) translated = domain.translate( - [itir.OffsetLiteral(value="V2E"), itir.OffsetLiteral(value=0)], offset_provider + [itir.OffsetLiteral(value=V2EDim.tag), itir.OffsetLiteral(value=0)], offset_provider ) assert translated.empty() assert set(translated.ranges.keys()) == {Edge} @@ -242,19 +260,19 @@ def test_is_finite_symbolic_domain(ranges, expected): @pytest.mark.parametrize( "shift_chain, expected_end_domain", [ - (("V2V", 0), {Vertex: (0, 4)}), - (("V2V", 1), {Vertex: (0, 4)}), - (("V2V", 2), {Vertex: (0, 1)}), - (("V2V", 3), {Vertex: (1, 4)}), - (("V2V", 0, "V2V", 3, "V2V", 0), {Vertex: (1, 4)}), - (("V2E", 0), {Edge: (0, 4)}), - (("V2E", 0, "E2V", 0), {Vertex: (0, 4)}), - (("V2V", 3, "V2E", 0), {Edge: (1, 4)}), + ((V2VDim.tag, 0), {Vertex: (0, 4)}), + ((V2VDim.tag, 1), {Vertex: (0, 4)}), + ((V2VDim.tag, 2), {Vertex: (0, 1)}), + ((V2VDim.tag, 3), {Vertex: (1, 4)}), + ((V2VDim.tag, 0, V2VDim.tag, 3, V2VDim.tag, 0), {Vertex: (1, 4)}), + ((V2EDim.tag, 0), {Edge: (0, 4)}), + ((V2EDim.tag, 0, E2VDim.tag, 0), {Vertex: (0, 4)}), + ((V2VDim.tag, 3, V2EDim.tag, 0), {Edge: (1, 4)}), ], ) def test_unstructured_translate(shift_chain, expected_end_domain): offset_provider = { - "V2V": constructors.as_connectivity( + V2VDim.tag: constructors.as_connectivity( domain={Vertex: (0, 4), V2VDim: 5}, codomain=Vertex, data=np.asarray( @@ -262,7 +280,7 @@ def test_unstructured_translate(shift_chain, expected_end_domain): dtype=fbuiltins.IndexType, ), ), - "V2E": constructors.as_connectivity( + V2EDim.tag: constructors.as_connectivity( domain={Vertex: (0, 4), V2EDim: 1}, codomain=Edge, data=np.asarray( @@ -272,7 +290,7 @@ def test_unstructured_translate(shift_chain, expected_end_domain): dtype=fbuiltins.IndexType, ).reshape((4, 1)), ), - "E2V": constructors.as_connectivity( + E2VDim.tag: constructors.as_connectivity( domain={Edge: (0, 4), E2VDim: 1}, codomain=Vertex, data=np.asarray( @@ -299,7 +317,7 @@ def test_unstructured_translate_with_symbolic_domain_sizes(as_type): # 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`. offset_provider = { - "V2E": constructors.as_connectivity( + V2EDim.tag: constructors.as_connectivity( domain={Vertex: (0, 4), V2EDim: 1}, codomain=Edge, data=np.asarray([0, 1, 2, 3], dtype=fbuiltins.IndexType).reshape((4, 1)), @@ -312,9 +330,9 @@ def test_unstructured_translate_with_symbolic_domain_sizes(as_type): im.domain(common.GridType.UNSTRUCTURED, {Vertex: (0, 4)}) ) translated = domain.translate( - [itir.OffsetLiteral(value="V2E"), itir.OffsetLiteral(value=0)], + [itir.OffsetLiteral(value=V2EDim.tag), itir.OffsetLiteral(value=0)], offset_provider, - symbolic_domain_sizes={"Edge": im.ref("num_edges")}, + symbolic_domain_sizes={Edge.tag: im.ref("num_edges")}, ) expected = im.domain(common.GridType.UNSTRUCTURED, {Edge: (0, im.ref("num_edges"))}) @@ -338,13 +356,13 @@ def test_non_contiguous_domain_warning(monkeypatch): monkeypatch.setattr(domain_utils, "_NON_CONTIGUOUS_DOMAIN_WARNING_SKIPPED_OFFSET_TAGS", set()) offset_provider = { - "V2V": constructors.as_connectivity( + V2VDim.tag: constructors.as_connectivity( domain={Vertex: (0, 100), V2VDim: 1}, codomain=Vertex, data=np.asarray([0] + [99] * 99, dtype=fbuiltins.IndexType).reshape((100, 1)), ) } - shift_chain = ("V2V", 0) + shift_chain = (V2VDim.tag, 0) shift_chain = [im.ensure_offset(o) for o in shift_chain] domain = domain_utils.SymbolicDomain.from_expr( im.domain(common.GridType.UNSTRUCTURED, {Vertex: (0, 2)}) @@ -358,13 +376,13 @@ def test_non_contiguous_domain_warning(monkeypatch): def test_oob_error(): offset_provider = { - "V2V": constructors.as_connectivity( + V2VDim.tag: constructors.as_connectivity( domain={Vertex: (0, 3), V2VDim: 1}, codomain=Vertex, data=np.asarray([0, -1, 1], dtype=fbuiltins.IndexType).reshape((3, 1)), ) } - shift_chain = ("V2V", 0) + shift_chain = (V2VDim.tag, 0) shift_chain = [im.ensure_offset(o) for o in shift_chain] domain = domain_utils.SymbolicDomain.from_expr( im.domain(common.GridType.UNSTRUCTURED, {Vertex: (0, 3)}) 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 3a6562e6be..12c0b75649 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 @@ -23,10 +23,16 @@ ) -E = gtx.Dimension("E") -V = gtx.Dimension("V") -E2VDim = gtx.Dimension("E2V", kind=gtx.DimensionKind.LOCAL) -E2V = gtx.FieldOffset("E2V", source=V, target=(E, E2VDim)) +class E(gtx.DimensionIndex): ... + + +class V(gtx.DimensionIndex): ... + + +class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +E2V = gtx.FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) # 0 --0-- 1 --1-- 2 @@ -44,7 +50,7 @@ def testee(inp): return as_fieldop(lambda it: neighbors(E2V, it), domain)(inp) inp = gtx.as_field([V], np.arange(3)) - with embedded_context.update(offset_provider={"E2V": e2v_conn}): + with embedded_context.update(offset_provider={E2VDim.tag: e2v_conn}): result = testee(inp) ref = e2v_arr @@ -76,7 +82,7 @@ def testee(inp): ) inp = gtx.as_field([V], np.arange(3)) - with embedded_context.update(offset_provider={"E2V": e2v_conn}): + with embedded_context.update(offset_provider={E2VDim.tag: e2v_conn}): result = testee(inp) ref = e2v_arr + 42.0 @@ -94,7 +100,7 @@ def testee(inp, mask): inp = gtx.as_field([V], np.arange(3)) mask_field = gtx.as_field([E], np.array([True, False])) - with embedded_context.update(offset_provider={"E2V": e2v_conn}): + with embedded_context.update(offset_provider={E2VDim.tag: e2v_conn}): result = testee(inp, mask_field) ref = np.empty_like(e2v_arr, dtype=float) @@ -122,7 +128,7 @@ def testee(inp, mask): inp = gtx.as_field([V], np.arange(3)) mask_field = gtx.as_field([E], np.array([True, False])) - with embedded_context.update(offset_provider={"E2V": e2v_conn}): + with embedded_context.update(offset_provider={E2VDim.tag: e2v_conn}): result = testee(inp, mask_field) ref = np.empty_like(e2v_arr, dtype=float) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py index 9f61e5317e..894a85fa1d 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py @@ -17,6 +17,9 @@ from gt4py.next.iterator import embedded +class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + + def test_column_ufunc(): def test_func(): a = embedded.Column(1, np.asarray(range(0, 3))) @@ -44,9 +47,7 @@ def test_func(data_a: int, data_b: int): with embedded_context.update( offset_provider={}, - closure_column_range=common.NamedRange( - common.Dimension("K", kind=common.DimensionKind.VERTICAL), range(0, 3) - ), + closure_column_range=common.NamedRange(K, range(0, 3)), ): test_func(2, 3) @@ -148,6 +149,5 @@ def test_func(): def test_lift_accepts_cartesian_dimension_offset(): - K = common.Dimension("K", kind=common.DimensionKind.VERTICAL) lifted = embedded.lift(lambda *args: 0)() lifted.shift(common.CartesianConnectivity(K), 1) # must not raise diff --git a/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py b/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py index fb159de5bb..b5432098bc 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py @@ -11,7 +11,10 @@ from gt4py.next.iterator.transforms import inline_dynamic_shifts from gt4py.next.type_system import type_specifications as ts -IDim = gtx.Dimension("IDim") + +class IDim(gtx.DimensionIndex): ... + + field_type = ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.INT32)) IOff = im.cartesian_offset(IDim, IDim) 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 bf2df06bf2..3cba7b16cd 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 @@ -15,19 +15,29 @@ from gt4py.next.iterator.runtime import CartesianDomain, UnstructuredDomain, _deduce_domain, fundef +class dummy_codomain(gtx.DimensionIndex): ... + + +class dummy_origin(gtx.DimensionIndex): ... + + +class dummy_neighbor(gtx.DimensionIndex): ... + + @fundef def foo(inp): return deref(inp) connectivity = common.ConnectivityType( - domain=[gtx.Dimension("dummy_origin"), gtx.Dimension("dummy_neighbor")], - codomain=gtx.Dimension("dummy_codomain"), + domain=[dummy_origin, dummy_neighbor], + codomain=dummy_codomain, skip_value=common._DEFAULT_SKIP_VALUE, dtype=None, ) -I = gtx.Dimension("I") + +class I(gtx.DimensionIndex): ... def test_deduce_domain(): diff --git a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py index 92bc96e714..94f0610c06 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py @@ -26,6 +26,7 @@ from gt4py.next.type_system import type_specifications as ts from next_tests.integration_tests.cases import ( + C2EDim, C2E, E2V, V2E, @@ -101,18 +102,18 @@ def expression_test_cases(): ), ( im.named_range( - itir.AxisLiteral(value="Vertex", kind=common.DimensionKind.HORIZONTAL), 0, 1 + itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 ), it_ts.NamedRangeType(dim=Vertex), ), ( - im.call("cartesian_domain")(im.named_range(itir.AxisLiteral(value="IDim"), 0, 1)), + im.call("cartesian_domain")(im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1)), ts.DomainType(dims=[IDim]), ), ( im.call("unstructured_domain")( im.named_range( - itir.AxisLiteral(value="Vertex", kind=common.DimensionKind.HORIZONTAL), 0, 1 + itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 ) ), ts.DomainType(dims=[Vertex]), @@ -131,7 +132,7 @@ def expression_test_cases(): ), # neighbors ( - im.neighbors("E2V", im.ref("a", it_on_e_of_e_type)), + im.neighbors(E2VDim.tag, im.ref("a", it_on_e_of_e_type)), ts.ListType(element_type=it_on_e_of_e_type.element_type, offset_type=E2VDim), ), # cast @@ -183,7 +184,7 @@ def expression_test_cases(): ts.TupleType(types=[int_type, float64_type]), ), # shift - (im.shift("V2E", 1)(im.ref("it", it_on_v_of_e_type)), it_on_e_of_e_type), + (im.shift(V2EDim.tag, 1)(im.ref("it", it_on_v_of_e_type)), it_on_e_of_e_type), # cartesian shift via `CartesianOffset` (im.shift(Ioff, 1)(im.ref("it", it_ijk_type)), it_ijk_type), # as_fieldop @@ -207,7 +208,7 @@ def expression_test_cases(): ), ( im.as_fieldop( - im.lambda_("it")(im.deref(im.shift(Koff, 1)(im.shift("V2E", 0)("it")))), + im.lambda_("it")(im.deref(im.shift(Koff, 1)(im.shift(V2EDim.tag, 0)("it")))), vertex_k_domain, )(im.ref("inp", float_edge_k_field)), float_vertex_k_field, @@ -216,7 +217,7 @@ def expression_test_cases(): im.as_fieldop( im.lambda_("it1", "it2")( im.plus( - im.deref(im.shift("E2V", 1)(im.shift("C2E", 1)("it1"))), + im.deref(im.shift(E2VDim.tag, 1)(im.shift(C2EDim.tag, 1)("it1"))), im.deref(im.shift(Koff, 1)("it2")), ), ), @@ -229,7 +230,7 @@ def expression_test_cases(): ), ( im.as_fieldop( - im.lambda_("it")(im.deref(im.shift("V2E", 0)("it"))), + im.lambda_("it")(im.deref(im.shift(V2EDim.tag, 0)("it"))), vertex_k_domain, )(im.ref("inp", float_edge_k_field)), float_vertex_k_field, @@ -261,13 +262,13 @@ def expression_test_cases(): im.as_fieldop( im.lambda_("a", "b")(im.plus(im.deref("a"), im.deref("b"))), im.call("cartesian_domain")( - im.named_range(itir.AxisLiteral(value="IDim"), 0, 1) + im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1) ), )(im.ref("inp", float_i_field), 1.0), im.as_fieldop( "deref", im.call("cartesian_domain")( - im.named_range(itir.AxisLiteral(value="IDim"), 0, 1) + im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1) ), )(im.ref("inp", float_i_field)), ), @@ -389,7 +390,7 @@ def test_late_offset_axis(): mesh = simple_mesh(None) func = im.lambda_("dim")(im.shift(im.ref("dim"), 1)(im.ref("it", it_on_v_of_e_type))) - testee = im.call(func)(im.ensure_offset("V2E")) + testee = im.call(func)(im.ensure_offset(V2EDim.tag)) result = itir_type_inference.infer( testee, offset_provider_type=mesh.offset_provider_type, allow_undeclared_symbols=True @@ -412,7 +413,7 @@ def test_cast_first_arg_inference(): def test_cartesian_fencil_definition(): cartesian_domain = im.call("cartesian_domain")( - im.named_range(itir.AxisLiteral(value="IDim"), 0, 1) + im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1) ) testee = itir.Program( @@ -443,9 +444,9 @@ def test_unstructured_fencil_definition(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( im.named_range( - itir.AxisLiteral(value="Vertex", kind=common.DimensionKind.HORIZONTAL), 0, 1 + itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 ), - im.named_range(itir.AxisLiteral(value="KDim", kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), ) testee = itir.Program( @@ -456,7 +457,7 @@ def test_unstructured_fencil_definition(): body=[ itir.SetAt( expr=im.as_fieldop( - im.lambda_("it")(im.deref(im.shift("V2E", 0)("it"))), unstructured_domain + im.lambda_("it")(im.deref(im.shift(V2EDim.tag, 0)("it"))), unstructured_domain )(im.ref("inp")), domain=unstructured_domain, target=im.ref("out"), @@ -478,7 +479,7 @@ def test_unstructured_fencil_definition(): def test_function_definition(): cartesian_domain = im.call("cartesian_domain")( - im.named_range(itir.AxisLiteral(value="IDim"), 0, 1) + im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1) ) testee = itir.Program( @@ -510,9 +511,9 @@ def test_fencil_with_nb_field_input(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( im.named_range( - itir.AxisLiteral(value="Vertex", kind=common.DimensionKind.HORIZONTAL), 0, 1 + itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 ), - im.named_range(itir.AxisLiteral(value="KDim", kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), ) testee = itir.Program( @@ -540,7 +541,7 @@ def test_fencil_with_nb_field_input(): def test_program_tuple_setat_short_target(): cartesian_domain = im.call("cartesian_domain")( - im.named_range(itir.AxisLiteral(value="IDim"), 0, 1) + im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1) ) testee = itir.Program( @@ -571,7 +572,7 @@ def test_program_tuple_setat_short_target(): def test_program_setat_without_domain(): cartesian_domain = im.call("cartesian_domain")( - im.named_range(itir.AxisLiteral(value="IDim"), 0, 1) + im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1) ) testee = itir.Program( @@ -595,7 +596,7 @@ def test_program_setat_without_domain(): def test_if_stmt(): cartesian_domain = im.call("cartesian_domain")( - im.named_range(itir.AxisLiteral(value="IDim"), 0, 1) + im.named_range(itir.AxisLiteral(value=IDim.tag), 0, 1) ) testee = itir.IfStmt( @@ -622,7 +623,7 @@ def test_as_fieldop_without_domain_nb_field_input(): testee = im.as_fieldop(stencil)(im.ref("inp1", float_vertex_v2e_field)) result = itir_type_inference.infer( - testee, offset_provider_type={"V2E": V2E}, allow_undeclared_symbols=True + testee, offset_provider_type={V2EDim.tag: V2E}, allow_undeclared_symbols=True ) assert result.type == ts.FieldType(dims=[Vertex], dtype=float64_list_type) assert result.fun.args[0].type.pos_only_args[0] == it_ts.IteratorType( diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_collapse_tuple.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_collapse_tuple.py index 18d1bdc08c..f69fbc768e 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_collapse_tuple.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_collapse_tuple.py @@ -15,7 +15,9 @@ int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) -Vertex = common.Dimension(value="Vertex", kind=common.DimensionKind.HORIZONTAL) + + +class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... def test_simple_make_tuple_tuple_get(uids: utils.IDGeneratorPool): diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py index 58086670f9..b6a7d05388 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py @@ -16,7 +16,11 @@ from gt4py.next.type_system import type_specifications as ts int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) -IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL) + + +class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + field_type = ts.FieldType(dims=[IDim], dtype=int_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py index efb050d516..97c2f4ce40 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py @@ -17,7 +17,11 @@ from gt4py.next.iterator.type_system import type_specifications as it_ts int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) -IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL) + + +class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + field_type = ts.FieldType(dims=[IDim], dtype=int_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py index 2517c7ab55..babcafc4a1 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py @@ -14,8 +14,12 @@ from gt4py.next.type_system import type_specifications as ts int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) -IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL) -JDim = common.Dimension(value="JDim", kind=common.DimensionKind.HORIZONTAL) + + +class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class JDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... def test_in_helper(): diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py index b5551b6af5..e2e73a819f 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py @@ -20,9 +20,12 @@ ) +class I(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + @pytest.fixture def offset_provider_type(request): - return {"I": common.Dimension("I", kind=common.DimensionKind.HORIZONTAL)} + return {"I": I} @pytest.fixture diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py index 6f519b246b..f2e4f4bb07 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py @@ -14,7 +14,10 @@ from gt4py.next.iterator import ir as itir from gt4py.next.iterator.transforms import dead_code_elimination -TDim = common.Dimension(value="TDim") + +class TDim(common.DimensionIndex): ... + + int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) field_type = ts.FieldType(dims=[TDim], dtype=int_type) 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 317ce27774..eafe45b9d3 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 @@ -28,12 +28,26 @@ float_type = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) -IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL) -JDim = common.Dimension(value="JDim", kind=common.DimensionKind.HORIZONTAL) -KDim = common.Dimension(value="KDim", kind=common.DimensionKind.VERTICAL) -Vertex = common.Dimension(value="Vertex", kind=common.DimensionKind.HORIZONTAL) -Edge = common.Dimension(value="Edge", kind=common.DimensionKind.HORIZONTAL) -E2VDim = common.Dimension(value="E2V", kind=common.DimensionKind.LOCAL) + + +class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class JDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + + +class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + float_i_field = ts.FieldType(dims=[IDim], dtype=float_type) float_ij_field = ts.FieldType(dims=[IDim, JDim], dtype=float_type) tuple_float_i_field = ts.TupleType( @@ -47,7 +61,7 @@ @pytest.fixture def unstructured_offset_provider(): return { - "E2V": constructors.as_connectivity( + E2VDim.tag: constructors.as_connectivity( domain={Edge: 1, E2VDim: 2}, codomain=Vertex, data=np.array([[0, 1]], dtype=np.int32), @@ -223,9 +237,9 @@ def test_multi_length_shift(): def test_unstructured_shift(unstructured_offset_provider): - stencil = im.lambda_("arg0")(im.deref(im.shift("E2V", 1)("arg0"))) + stencil = im.lambda_("arg0")(im.deref(im.shift(E2VDim.tag, 1)("arg0"))) domain = im.domain(common.GridType.UNSTRUCTURED, {Edge: (0, 1)}) - accessed_vertex = unstructured_offset_provider["E2V"].ndarray[0, 1] + accessed_vertex = unstructured_offset_provider[E2VDim.tag].ndarray[0, 1] expected_domains = {"in_field1": {Vertex: (accessed_vertex, accessed_vertex + np.int32(1))}} testee, expected = setup_test_as_fieldop(stencil, domain, expected_domains=expected_domains) @@ -1168,9 +1182,9 @@ def test_scan(): def test_symbolic_domain_sizes(unstructured_offset_provider): - stencil = im.lambda_("arg0")(im.deref(im.shift("E2V", 1)("arg0"))) + stencil = im.lambda_("arg0")(im.deref(im.shift(E2VDim.tag, 1)("arg0"))) domain = im.domain(common.GridType.UNSTRUCTURED, {Edge: (0, 1)}) - symbolic_domain_sizes = {"Vertex": "num_vertices"} + symbolic_domain_sizes = {Vertex.tag: "num_vertices"} expected_domains = {"in_field1": {Vertex: (0, im.ref("num_vertices"))}} testee, expected = setup_test_as_fieldop( stencil, @@ -1412,7 +1426,7 @@ def test_concat_where_unstructured_shift_in_never_selected_branch(unstructured_o # `Edge: [1, 1)` range must not raise. domain = im.domain(common.GridType.UNSTRUCTURED, {Edge: (0, 1)}) cond = im.domain(common.GridType.UNSTRUCTURED, {Edge: (1, itir.InfinityLiteral.POSITIVE)}) - stencil_e2v = im.lambda_("it")(im.deref(im.shift("E2V", 0)("it"))) + stencil_e2v = im.lambda_("it")(im.deref(im.shift(E2VDim.tag, 0)("it"))) domain_empty = im.domain(common.GridType.UNSTRUCTURED, {Edge: (1, 1)}) domain_a = im.domain(common.GridType.UNSTRUCTURED, {Vertex: (0, 0)}) @@ -1439,7 +1453,7 @@ def test_concat_where_unstructured_shift_in_never_selected_branch(unstructured_o def test_broadcast(): - testee = im.call("broadcast")("in_field", im.make_tuple(itir.AxisLiteral(value="IDim"))) + testee = im.call("broadcast")("in_field", im.make_tuple(itir.AxisLiteral(value=IDim.tag))) domain = im.domain(common.GridType.CARTESIAN, {IDim: (0, 10)}) expected_domains = { "in_field": {IDim: (0, 10)}, diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py index 4dc3373832..d8f3327386 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py @@ -16,7 +16,9 @@ from gt4py.next.type_system import type_specifications as ts -IDim = common.Dimension("IDim") +class IDim(common.DimensionIndex): ... + + T = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) i_field = ts.FieldType(dims=[IDim], dtype=T) i_tuple_field = ts.TupleType(types=[i_field, i_field]) 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 543df6beb8..5bceba0e53 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,8 +17,15 @@ from gt4py.next.type_system import type_specifications as ts -IDim = common.Dimension("IDim") -JDim = common.Dimension("JDim") +class Neighbor(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class IDim(common.DimensionIndex): ... + + +class JDim(common.DimensionIndex): ... + + field_type = ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.INT32)) IOff = im.cartesian_offset(IDim, IDim) @@ -356,7 +363,7 @@ def test_inline_as_fieldop_with_list_dtype(uids: utils.IDGeneratorPool): dims=[IDim], dtype=ts.ListType( element_type=ts.ScalarType(kind=ts.ScalarKind.INT32), - offset_type=common.Dimension("Neighbor", kind=common.DimensionKind.LOCAL), + offset_type=Neighbor, ), ) d = im.domain("cartesian_domain", {IDim: (0, 1)}) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py index 14e626be3a..9adfdaedbd 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py @@ -26,9 +26,15 @@ ) -IDim = common.Dimension(value="IDim") -JDim = common.Dimension(value="JDim") -KDim = common.Dimension(value="KDim", kind=common.DimensionKind.VERTICAL) +class IDim(common.DimensionIndex): ... + + +class JDim(common.DimensionIndex): ... + + +class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + + index_type = ts.ScalarType(kind=getattr(ts.ScalarKind, builtins.INTEGER_INDEX_BUILTIN.upper())) float_type = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) i_field_type = ts.FieldType(dims=[IDim], dtype=float_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py index 4569b42305..3a7c92b784 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py @@ -13,7 +13,10 @@ from gt4py.next.iterator.transforms import inline_scalar from gt4py.next.iterator.ir_utils import ir_makers as im -TDim = common.Dimension(value="TDim") + +class TDim(common.DimensionIndex): ... + + int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py index b1a18ddab8..26cff82c6c 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py @@ -13,6 +13,9 @@ from gt4py.next.type_system import type_specifications as ts +class IDim(gtx.DimensionIndex): ... + + def test_prune_casts_simple(): x_ref = im.ref("x", ts.ScalarType(kind=ts.ScalarKind.FLOAT32)) y_ref = im.ref("y", ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) @@ -25,7 +28,6 @@ def test_prune_casts_simple(): def test_prune_casts_fieldop(): - IDim = gtx.Dimension("IDim") x_ref = im.ref("x", ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT32))) y_ref = im.ref("y", ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64))) testee = im.op_as_fieldop("plus")( 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 ead2bd1724..3f6cd212ae 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 @@ -18,10 +18,18 @@ from gt4py.next.iterator.ir_utils import common_pattern_matcher as cpm, domain_utils from gt4py.next.type_system import type_info, type_specifications as ts -Vertex = common.Dimension(value="Vertex", kind=common.DimensionKind.HORIZONTAL) -Edge = common.Dimension(value="Edge", kind=common.DimensionKind.HORIZONTAL) -V2EDim = common.Dimension(value="V2E", kind=common.DimensionKind.LOCAL) -K = common.Dimension(value="K", kind=common.DimensionKind.VERTICAL) + +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 K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + float64 = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) vertex_k_field = ts.FieldType(dims=[Vertex, K], dtype=float64) @@ -246,13 +254,13 @@ def test_prune_equal_branches_containing_unstructured_shift(): failing on the `V2E` translation during the re-inference). """ offset_provider = { - "V2E": constructors.as_connectivity( + V2EDim.tag: constructors.as_connectivity( domain={Vertex: 1, V2EDim: 2}, codomain=Edge, data=np.array([[0, 1]], dtype=np.int32), ) } - stencil = im.lambda_("it")(im.deref(im.shift("V2E", 0)("it"))) + stencil = im.lambda_("it")(im.deref(im.shift(V2EDim.tag, 0)("it"))) branch = im.as_fieldop(stencil)(im.ref("e", edge_field)) accessed_domain = {Vertex: (0, 1), K: (0, 10)} 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 8083c56a8f..5451a2dc44 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 @@ -15,25 +15,40 @@ from gt4py.next.type_system import type_specifications as ts +class dummy_codomain(common.DimensionIndex): ... + + +class dummy_origin(common.DimensionIndex): ... + + +class dummy_neighbor(common.DimensionIndex): ... + + +#: The local dimensions of the neighbor lists under test. Each one's `tag` is also its IR offset +#: string and its offset-provider key: `UnrollReduce` looks a connectivity up by the local +#: dimension of the list it reduces, so those three names must be a single string (ADR 0028). +class Dim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class Dim2(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + def dummy_connectivity_type(max_neighbors: int, has_skip_values: bool): return common.NeighborConnectivityType( - domain=[common.Dimension("dummy_origin"), common.Dimension("dummy_neighbor")], - codomain=common.Dimension("dummy_codomain"), + domain=[dummy_origin, dummy_neighbor], + codomain=dummy_codomain, skip_value=common._DEFAULT_SKIP_VALUE if has_skip_values else None, dtype=None, max_neighbors=max_neighbors, ) -def _list_type(dim: str) -> ts.ListType: - return ts.ListType( - element_type=ts.DataType(), - offset_type=common.Dimension(value=dim, kind=common.DimensionKind.LOCAL), - ) +def _list_type(dim: common.Dimension) -> ts.ListType: + return ts.ListType(element_type=ts.DataType(), offset_type=dim) -def typed_neighbors(dim: str, arg: str | ir.Expr) -> ir.FunCall: - neighbors = im.neighbors(dim, arg) +def typed_neighbors(dim: common.Dimension, arg: str | ir.Expr) -> ir.FunCall: + neighbors = im.neighbors(dim.tag, arg) neighbors.type = _list_type(dim) return neighbors @@ -45,32 +60,32 @@ def has_skip_values(request): @pytest.fixture def basic_reduction(): - return im.reduce("foo", 0.0)(typed_neighbors("Dim", "x")) + return im.reduce("foo", 0.0)(typed_neighbors(Dim, "x")) @pytest.fixture def reduction_with_shift_on_second_arg(): const_list = im.call("make_const_list")(42) - const_list.type = _list_type("Dim") - return im.reduce("foo", 0.0)(const_list, typed_neighbors("Dim", "y")) + const_list.type = _list_type(Dim) + return im.reduce("foo", 0.0)(const_list, typed_neighbors(Dim, "y")) @pytest.fixture def reduction_with_incompatible_shifts(): - return im.reduce("foo", 0.0)(typed_neighbors("Dim", "x"), typed_neighbors("Dim2", "y")) + return im.reduce("foo", 0.0)(typed_neighbors(Dim, "x"), typed_neighbors(Dim2, "y")) @pytest.fixture def reduction_with_irrelevant_full_shift(): return im.reduce("foo", 0.0)( - typed_neighbors("Dim", im.shift("IrrelevantDim", 0)("x")), typed_neighbors("Dim", "y") + typed_neighbors(Dim, im.shift("IrrelevantDim", 0)("x")), typed_neighbors(Dim, "y") ) @pytest.fixture def reduction_if(): - if_expr = im.if_(True, typed_neighbors("Dim", "x"), "y") - if_expr.type = _list_type("Dim") + if_expr = im.if_(True, typed_neighbors(Dim, "x"), "y") + if_expr.type = _list_type(Dim) return im.reduce("foo", 0.0)(if_expr) @@ -86,7 +101,7 @@ def reduction_if(): def test_get_partial_offsets(reduction, request): partial_offsets = _get_partial_offset_tags(request.getfixturevalue(reduction).args) - assert set(partial_offsets) == {"Dim"} + assert set(partial_offsets) == {Dim.tag} def _expected(red, max_neighbors, has_skip_values, shifted_arg=0): @@ -116,7 +131,7 @@ def test_basic(basic_reduction, has_skip_values, uids: utils.IDGeneratorPool): expected = _expected(basic_reduction, 3, has_skip_values) offset_provider_type = { - "Dim": dummy_connectivity_type(max_neighbors=3, has_skip_values=has_skip_values) + Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=has_skip_values) } actual = UnrollReduce.apply( basic_reduction, offset_provider_type=offset_provider_type, uids=uids @@ -130,7 +145,7 @@ def test_reduction_with_shift_on_second_arg( expected = _expected(reduction_with_shift_on_second_arg, 1, has_skip_values, 1) offset_provider_type = { - "Dim": dummy_connectivity_type(max_neighbors=1, has_skip_values=has_skip_values) + Dim.tag: dummy_connectivity_type(max_neighbors=1, has_skip_values=has_skip_values) } actual = UnrollReduce.apply( reduction_with_shift_on_second_arg, offset_provider_type=offset_provider_type, uids=uids @@ -141,7 +156,9 @@ def test_reduction_with_shift_on_second_arg( def test_reduction_with_if(reduction_if, uids: utils.IDGeneratorPool): expected = _expected(reduction_if, 2, False) - offset_provider_type = {"Dim": dummy_connectivity_type(max_neighbors=2, has_skip_values=False)} + offset_provider_type = { + Dim.tag: dummy_connectivity_type(max_neighbors=2, has_skip_values=False) + } actual = UnrollReduce.apply(reduction_if, offset_provider_type=offset_provider_type, uids=uids) assert actual == expected @@ -152,7 +169,7 @@ def test_reduction_with_irrelevant_full_shift( expected = _expected(reduction_with_irrelevant_full_shift, 3, False) offset_provider_type = { - "Dim": dummy_connectivity_type(max_neighbors=3, has_skip_values=False), + Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=False), "IrrelevantDim": dummy_connectivity_type( max_neighbors=1, has_skip_values=True ), # different max_neighbors and skip value to trigger error @@ -167,16 +184,16 @@ def test_reduction_with_irrelevant_full_shift( "offset_provider_type", [ { - "Dim": dummy_connectivity_type(max_neighbors=3, has_skip_values=False), - "Dim2": dummy_connectivity_type(max_neighbors=2, has_skip_values=False), + Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=False), + Dim2.tag: dummy_connectivity_type(max_neighbors=2, has_skip_values=False), }, { - "Dim": dummy_connectivity_type(max_neighbors=3, has_skip_values=False), - "Dim2": dummy_connectivity_type(max_neighbors=3, has_skip_values=True), + Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=False), + Dim2.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=True), }, { - "Dim": dummy_connectivity_type(max_neighbors=3, has_skip_values=False), - "Dim2": dummy_connectivity_type(max_neighbors=2, has_skip_values=True), + Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=False), + Dim2.tag: dummy_connectivity_type(max_neighbors=2, has_skip_values=True), }, ], ) diff --git a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py index a25732649a..371d7730de 100644 --- a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py +++ b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py @@ -15,6 +15,12 @@ from gt4py.next.otf.binding import interface +class bar(gtx.DimensionIndex): ... + + +class foo(gtx.DimensionIndex): ... + + @pytest.fixture def function_scalar_example(): return interface.Function( @@ -60,15 +66,13 @@ def function_buffer_example(): interface.Parameter( name="a_buf", type_=ts.FieldType( - dims=[gtx.Dimension("bar"), gtx.Dimension("foo")], + dims=[bar, foo], dtype=ts.ScalarType(ts.ScalarKind.FLOAT64), ), ), interface.Parameter( name="b_buf", - type_=ts.FieldType( - dims=[gtx.Dimension("bar")], dtype=ts.ScalarType(ts.ScalarKind.INT64) - ), + type_=ts.FieldType(dims=[bar], dtype=ts.ScalarType(ts.ScalarKind.INT64)), ), ], ) @@ -111,11 +115,11 @@ def function_tuple_example(): type_=ts.TupleType( types=[ ts.FieldType( - dims=[gtx.Dimension("bar"), gtx.Dimension("foo")], + dims=[bar, foo], dtype=ts.ScalarType(ts.ScalarKind.FLOAT64), ), ts.FieldType( - dims=[gtx.Dimension("bar"), gtx.Dimension("foo")], + dims=[bar, foo], dtype=ts.ScalarType(ts.ScalarKind.FLOAT64), ), ] diff --git a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py index ba83a07b2d..00c77dfd58 100644 --- a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py +++ b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py @@ -13,12 +13,18 @@ import gt4py.next as gtx import gt4py.next.type_system.type_specifications as ts -from gt4py.next import config +from gt4py.next import common, config from gt4py.next.otf import artifacts from gt4py.next.otf.binding import cpp_interface, interface, nanobind from gt4py.next.otf.compilation import cache +class I(gtx.DimensionIndex): ... + + +class J(gtx.DimensionIndex): ... + + def make_program_source(name: str) -> artifacts.ProgramSource: entry_point = interface.Function( name, @@ -26,7 +32,7 @@ def make_program_source(name: str) -> artifacts.ProgramSource: interface.Parameter( name="buf", type_=ts.FieldType( - dims=[gtx.Dimension("I"), gtx.Dimension("J")], + dims=[I, J], dtype=ts.ScalarType(ts.ScalarKind.FLOAT32), ), ), @@ -35,11 +41,11 @@ def make_program_source(name: str) -> artifacts.ProgramSource: type_=ts.TupleType( types=[ ts.FieldType( - dims=[gtx.Dimension("I"), gtx.Dimension("J")], + dims=[I, J], dtype=ts.ScalarType(ts.ScalarKind.FLOAT32), ), ts.FieldType( - dims=[gtx.Dimension("I"), gtx.Dimension("J")], + dims=[I, J], dtype=ts.ScalarType(ts.ScalarKind.FLOAT32), ), ] @@ -49,11 +55,14 @@ def make_program_source(name: str) -> artifacts.ProgramSource: ), returns=True, ) + # NOTE: the tag types are named after the *mangled* dimension tags, which is what the + # generated bindings reference; a dimension's tag is its qualified Python name (ADR 0028). + i_t, j_t = (common.codegen_name(d.tag) for d in (I, J)) func = cpp_interface.render_function_declaration( entry_point, - """\ - const auto xdim = gridtools::at_key(sid_get_upper_bounds(buf)); - const auto ydim = gridtools::at_key(sid_get_upper_bounds(buf)); + f"""\ + const auto xdim = gridtools::at_key(sid_get_upper_bounds(buf)); + const auto ydim = gridtools::at_key(sid_get_upper_bounds(buf)); return xdim * ydim * sc;\ """, ) @@ -62,12 +71,12 @@ def make_program_source(name: str) -> artifacts.ProgramSource: #include #include namespace generated { - struct I_t {} constexpr inline I; - struct J_t {} constexpr inline J; + struct {{i_t}}_t {} constexpr inline {{i_t}}; + struct {{j_t}}_t {} constexpr inline {{j_t}}; } {{func}}\ """ - ).render(func=func) + ).render(func=func, i_t=i_t, j_t=j_t) return artifacts.ProgramSource( entry_point=entry_point, diff --git a/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py b/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py index 3d9fcaa834..bd1be39266 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py +++ b/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py @@ -62,7 +62,7 @@ def test_sanitize_static_args_wrong_type(): static_arg.validate("foo", ts.ScalarType(kind=ts.ScalarKind.INT32)) -TDim = gtx.Dimension("TDim") +class TDim(gtx.DimensionIndex): ... @pytest.fixture 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 dfe0d5b5f6..c3c6aa6c7b 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_runners.py +++ b/tests/next_tests/unit_tests/otf_tests/test_runners.py @@ -25,6 +25,15 @@ from next_tests.fixtures import compilation as fixtures_compilation +class Vertex(gtx.DimensionIndex): ... + + +class Edge(gtx.DimensionIndex): ... + + +class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + @pytest.fixture def process_runner(tmp_path): runner = runners.ProcessRunner(max_workers=1, shared_session_cache_dir=str(tmp_path)) @@ -147,12 +156,9 @@ def compile(self, program, compile_time_args): def test_offloaded_task_ships_connectivities_as_file_refs(): - Vertex = gtx.Dimension("Vertex") - Edge = gtx.Dimension("Edge") - V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) conn = gtx.as_connectivity([Vertex, V2EDim], Edge, np.array([[0, 1], [1, 2], [2, 0]])) compile_time_args = dataclasses.replace( - arguments.CompileTimeArgs.empty(), offset_provider={"V2E": conn} + arguments.CompileTimeArgs.empty(), offset_provider={V2EDim.tag: conn} ) backend = next_backend.Backend( name="test_backend", @@ -166,9 +172,9 @@ def test_offloaded_task_ships_connectivities_as_file_refs(): ) # without refs the original compilable is used as is - assert task.construct_compilable(False).args.offset_provider["V2E"] is conn + assert task.construct_compilable(False).args.offset_provider[V2EDim.tag] is conn shipped = task.construct_compilable(True) - ref = shipped.args.offset_provider["V2E"] + ref = shipped.args.offset_provider[V2EDim.tag] assert isinstance(ref, compilation_tasks._ConnectivityFileRef) # task preparation is pure: nothing is dumped until a runner ships the task assert id(conn) not in compilation_tasks._connectivity_files @@ -183,14 +189,11 @@ def test_offloaded_task_ships_connectivities_as_file_refs(): task2 = compilation_tasks.make_compilation_task( backend, definition_stage=None, compile_time_args=compile_time_args ) - pickle.dumps(task2.construct_compilable(True).args.offset_provider["V2E"]) + pickle.dumps(task2.construct_compilable(True).args.offset_provider[V2EDim.tag]) assert compilation_tasks._connectivity_files[id(conn)][1] == path def test_connectivity_file_registry_prunes_on_gc(): - Vertex = gtx.Dimension("Vertex") - Edge = gtx.Dimension("Edge") - V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) conn = gtx.as_connectivity([Vertex, V2EDim], Edge, np.array([[0, 1], [1, 2], [2, 0]])) compilation_tasks._dump_connectivity(conn) diff --git a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py index af958b84d9..d233f29a8d 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py +++ b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py @@ -30,9 +30,11 @@ ) +class IDim(gtx.DimensionIndex): ... + + @pytest.fixture def program_example(): - IDim = gtx.Dimension("I") params = [gtx.as_field([IDim], np.empty((1,), dtype=np.float32)), np.float32(3.14)] param_types = [type_translation.from_value(param) for param in params] @@ -42,7 +44,7 @@ def program_example(): itir.FunCall( fun=itir.SymRef(id="named_range"), args=[ - itir.AxisLiteral(value="I"), + itir.AxisLiteral(value=IDim.tag), im.literal("0", builtins.INTEGER_INDEX_BUILTIN), im.literal("10", builtins.INTEGER_INDEX_BUILTIN), ], diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py index d381e0ff51..faea3c7764 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py @@ -22,7 +22,7 @@ from gt4py.next.ffront.fbuiltins import where from next_tests.integration_tests import cases -from next_tests.integration_tests.cases import E2V +from next_tests.integration_tests.cases import E2VDim, E2V from next_tests.integration_tests.cases_utils import ( Edge, IDim, @@ -208,7 +208,7 @@ def verify_testee(): def test_dace_fastcall_with_connectivity(unstructured_case, monkeypatch): """Test reuse of SDFG arguments between program calls by means of SDFG fastcall API.""" - connectivity_E2V = unstructured_case.offset_provider["E2V"].asnumpy() + connectivity_E2V = unstructured_case.offset_provider[E2VDim.tag].asnumpy() @gtx.field_operator def testee(a: cases.VField) -> cases.EField: diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index df65b39c4c..922ffbf3ff 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -376,8 +376,8 @@ def testee(a: cases.VField, b: cases.VField): b = cases.allocate(test_case, testee, "b")() ref = np.sum( - np.sum(a.asnumpy()[offset_provider["E2V"].asnumpy()], axis=1, initial=0)[ - offset_provider["V2E"].asnumpy() + np.sum(a.asnumpy()[offset_provider[E2VDim.tag].asnumpy()], axis=1, initial=0)[ + offset_provider[V2EDim.tag].asnumpy() ], axis=1, ) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py index 5a68f44961..8e9d037aa6 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py @@ -27,6 +27,9 @@ from gt4py.next.type_system import type_specifications as ts from next_tests.integration_tests.cases_utils import ( + E2VDim, + C2VDim, + C2EDim, Cell, Edge, IDim, @@ -88,10 +91,10 @@ def allow_view_arguments(): def make_mesh_symbols(mesh: MeshDescriptor): - c2e_ndarray = mesh.offset_provider["C2E"].ndarray - c2v_ndarray = mesh.offset_provider["C2V"].ndarray - e2v_ndarray = mesh.offset_provider["E2V"].ndarray - v2e_ndarray = mesh.offset_provider["V2E"].ndarray + c2e_ndarray = mesh.offset_provider[C2EDim.tag].ndarray + c2v_ndarray = mesh.offset_provider[C2VDim.tag].ndarray + e2v_ndarray = mesh.offset_provider[E2VDim.tag].ndarray + v2e_ndarray = mesh.offset_provider[V2EDim.tag].ndarray return dict( __cells_Cell_range_0=0, __cells_Cell_range_1=mesh.num_cells, @@ -998,15 +1001,17 @@ def test_gtir_connectivity_shift(): # apply shift 2 times along different dimensions stencil1_inlined = im.as_fieldop( im.lambda_("it")( - im.deref(im.shift("C2E", C2E_neighbor_idx)(im.shift("E2V", E2V_neighbor_idx)("it"))) + im.deref( + im.shift(C2EDim.tag, C2E_neighbor_idx)(im.shift(E2VDim.tag, E2V_neighbor_idx)("it")) + ) ) )("ev_field") # fieldview flavor of the same stncil: create an intermediate temporary field stencil1_fieldview = im.as_fieldop( - im.lambda_("it")(im.deref(im.shift("E2V", E2V_neighbor_idx)("it"))) + im.lambda_("it")(im.deref(im.shift(E2VDim.tag, E2V_neighbor_idx)("it"))) )( - im.as_fieldop(im.lambda_("it")(im.deref(im.shift("C2E", C2E_neighbor_idx)("it"))))( + im.as_fieldop(im.lambda_("it")(im.deref(im.shift(C2EDim.tag, C2E_neighbor_idx)("it"))))( "ev_field" ) ) @@ -1017,9 +1022,9 @@ def test_gtir_connectivity_shift(): im.deref( im.call( im.call("shift")( - im.ensure_offset("E2V"), + im.ensure_offset(E2VDim.tag), im.ensure_offset(E2V_neighbor_idx), - im.ensure_offset("C2E"), + im.ensure_offset(C2EDim.tag), im.ensure_offset(C2E_neighbor_idx), ) )("it") @@ -1033,9 +1038,9 @@ def test_gtir_connectivity_shift(): im.deref( im.call( im.call("shift")( - im.ensure_offset("E2V"), + im.ensure_offset(E2VDim.tag), im.plus(im.deref("e2v_off"), 0), - im.ensure_offset("C2E"), + im.ensure_offset(C2EDim.tag), im.deref("c2e_off"), ) )("it") @@ -1049,9 +1054,9 @@ def test_gtir_connectivity_shift(): im.deref( im.call( im.call("shift")( - im.ensure_offset("E2V"), + im.ensure_offset(E2VDim.tag), im.deref("e2v_off"), - im.ensure_offset("C2E"), + im.ensure_offset(C2EDim.tag), im.deref("c2e_off"), ) )("it") @@ -1068,8 +1073,8 @@ def test_gtir_connectivity_shift(): CELL_OFFSET_FTYPE = ts.FieldType(dims=[Cell], dtype=SIZE_TYPE) EDGE_OFFSET_FTYPE = ts.FieldType(dims=[Edge], dtype=SIZE_TYPE) - connectivity_C2E = SIMPLE_MESH.offset_provider["C2E"] - connectivity_E2V = SIMPLE_MESH.offset_provider["E2V"] + connectivity_C2E = SIMPLE_MESH.offset_provider[C2EDim.tag] + connectivity_E2V = SIMPLE_MESH.offset_provider[E2VDim.tag] ev = np.random.rand(SIMPLE_MESH.num_edges, SIMPLE_MESH.num_vertices) ref = ev[connectivity_C2E.asnumpy()[:, C2E_neighbor_idx], :][ @@ -1153,11 +1158,11 @@ def test_gtir_connectivity_shift_chain(): gtir.SetAt( expr=im.as_fieldop( # let domain inference infer the domain here - im.lambda_("it")(im.deref(im.shift("E2V", E2V_neighbor_idx)("it"))), + im.lambda_("it")(im.deref(im.shift(E2VDim.tag, E2V_neighbor_idx)("it"))), )( im.as_fieldop( # let domain inference infer the domain here - im.lambda_("it")(im.deref(im.shift("V2E", V2E_neighbor_idx)("it"))), + im.lambda_("it")(im.deref(im.shift(V2EDim.tag, V2E_neighbor_idx)("it"))), )("edges") ), domain=im.get_field_domain( @@ -1172,8 +1177,8 @@ def test_gtir_connectivity_shift_chain(): sdfg = build_dace_sdfg(testee, SIMPLE_MESH.offset_provider) - connectivity_E2V = SIMPLE_MESH.offset_provider["E2V"] - connectivity_V2E = SIMPLE_MESH.offset_provider["V2E"] + connectivity_E2V = SIMPLE_MESH.offset_provider[E2VDim.tag] + connectivity_V2E = SIMPLE_MESH.offset_provider[V2EDim.tag] e = np.random.rand(SIMPLE_MESH.num_edges) ref = e[ @@ -1223,7 +1228,7 @@ def test_gtir_neighbors_as_input(): gtir.SetAt( expr=im.let( "x", - im.as_fieldop_neighbors("V2E", "edges", outer_domain), + im.as_fieldop_neighbors(V2EDim.tag, "edges", outer_domain), )( im.as_fieldop( im.lambda_("it")( @@ -1242,7 +1247,7 @@ def test_gtir_neighbors_as_input(): # based on canonical order of field dimensions sdfg = build_dace_sdfg(testee, SIMPLE_MESH.offset_provider, skip_domain_inference=True) - connectivity_V2E = SIMPLE_MESH.offset_provider["V2E"] + connectivity_V2E = SIMPLE_MESH.offset_provider[V2EDim.tag] v2e_field = np.random.rand(SIMPLE_MESH.num_vertices, connectivity_V2E.shape[1], MESH_NUM_LEVELS) e = np.random.rand(SIMPLE_MESH.num_edges, MESH_NUM_LEVELS) @@ -1295,14 +1300,14 @@ def test_gtir_reduce(): init_value = np.random.rand() stencil_inlined = im.as_fieldop( im.lambda_("it")( - im.reduce("plus", im.literal_from_value(init_value))(im.neighbors("V2E", "it")) + im.reduce("plus", im.literal_from_value(init_value))(im.neighbors(V2EDim.tag, "it")) ) )("edges") stencil_fieldview = im.as_fieldop( im.lambda_("it")(im.reduce("plus", im.literal_from_value(init_value))(im.deref("it"))) - )(im.as_fieldop_neighbors("V2E", "edges")) + )(im.as_fieldop_neighbors(V2EDim.tag, "edges")) - connectivity_V2E = SIMPLE_MESH.offset_provider["V2E"] + connectivity_V2E = SIMPLE_MESH.offset_provider[V2EDim.tag] e = np.random.rand(SIMPLE_MESH.num_edges) v_ref = [ @@ -1350,14 +1355,14 @@ def test_gtir_reduce_with_skip_values(): init_value = np.random.rand() stencil_inlined = im.as_fieldop( im.lambda_("it")( - im.reduce("plus", im.literal_from_value(init_value))(im.neighbors("V2E", "it")) + im.reduce("plus", im.literal_from_value(init_value))(im.neighbors(V2EDim.tag, "it")) ) )("edges") stencil_fieldview = im.as_fieldop( im.lambda_("it")(im.reduce("plus", im.literal_from_value(init_value))(im.deref("it"))) - )(im.as_fieldop_neighbors("V2E", "edges")) + )(im.as_fieldop_neighbors(V2EDim.tag, "edges")) - connectivity_V2E = SKIP_VALUE_MESH.offset_provider["V2E"] + connectivity_V2E = SKIP_VALUE_MESH.offset_provider[V2EDim.tag] e = np.random.rand(SKIP_VALUE_MESH.num_edges) v_ref = [ @@ -1406,7 +1411,7 @@ def test_gtir_reduce_with_skip_values(): def test_gtir_reduce_dot_product(): init_value = np.random.rand() - connectivity_V2E = SKIP_VALUE_MESH.offset_provider["V2E"] + connectivity_V2E = SKIP_VALUE_MESH.offset_provider[V2EDim.tag] v2e_field = np.random.rand(*connectivity_V2E.shape) e = np.random.rand(SKIP_VALUE_MESH.num_edges) @@ -1441,7 +1446,7 @@ def test_gtir_reduce_dot_product(): )( im.op_as_fieldop(im.map_list("plus"))( im.op_as_fieldop(im.map_list("multiplies"))( - im.as_fieldop_neighbors("V2E", "edges"), + im.as_fieldop_neighbors(V2EDim.tag, "edges"), "v2e_field", ), im.op_as_fieldop("make_const_list")(1.0), @@ -1492,7 +1497,7 @@ def test_gtir_reduce_with_cond_neighbors(use_sparse): im.if_( "pred", "v2e_field", - im.as_fieldop_neighbors("V2E", "edges"), + im.as_fieldop_neighbors(V2EDim.tag, "edges"), ) ), domain=im.get_field_domain(gtx_common.GridType.UNSTRUCTURED, "vertices", [Vertex]), @@ -1501,7 +1506,7 @@ def test_gtir_reduce_with_cond_neighbors(use_sparse): ], ) - connectivity_V2E = SKIP_VALUE_MESH.offset_provider["V2E"] + connectivity_V2E = SKIP_VALUE_MESH.offset_provider[V2EDim.tag] sdfg = build_dace_sdfg(testee, SKIP_VALUE_MESH.offset_provider) @@ -1744,8 +1749,8 @@ def test_gtir_let_lambda_with_connectivity(): C2E_neighbor_idx = 1 C2V_neighbor_idx = 2 - connectivity_C2E = SIMPLE_MESH.offset_provider["C2E"] - connectivity_C2V = SIMPLE_MESH.offset_provider["C2V"] + connectivity_C2E = SIMPLE_MESH.offset_provider[C2EDim.tag] + connectivity_C2V = SIMPLE_MESH.offset_provider[C2VDim.tag] testee = gtir.Program( id="let_lambda_with_connectivity", @@ -1761,13 +1766,13 @@ def test_gtir_let_lambda_with_connectivity(): expr=im.let( "x1", im.as_fieldop( - im.lambda_("it")(im.deref(im.shift("C2E", C2E_neighbor_idx)("it"))) + im.lambda_("it")(im.deref(im.shift(C2EDim.tag, C2E_neighbor_idx)("it"))) )("edges"), )( im.let( "x2", im.as_fieldop( - im.lambda_("it")(im.deref(im.shift("C2V", C2V_neighbor_idx)("it"))) + im.lambda_("it")(im.deref(im.shift(C2VDim.tag, C2V_neighbor_idx)("it"))) )("vertices"), )(im.op_as_fieldop("plus")("x1", "x2")) ), @@ -1818,7 +1823,7 @@ def test_gtir_let_lambda_with_origin(): gtir.SetAt( expr=im.let("e1", im.op_as_fieldop("plus")("edges", 1.0))( im.as_fieldop( - im.lambda_("it")(im.deref(im.shift("C2E", C2E_neighbor_idx)("it"))), + im.lambda_("it")(im.deref(im.shift(C2EDim.tag, C2E_neighbor_idx)("it"))), )("e1") ), domain=apply_margin_on_field_domain( @@ -1835,7 +1840,7 @@ def test_gtir_let_lambda_with_origin(): c = np.random.rand(SIMPLE_MESH.num_cells, MESH_NUM_LEVELS) e = np.random.rand(SIMPLE_MESH.num_edges, MESH_NUM_LEVELS) - connectivity_C2E = SIMPLE_MESH.offset_provider["C2E"] + connectivity_C2E = SIMPLE_MESH.offset_provider[C2EDim.tag] ref = np.concatenate( (c[:, :1], e[connectivity_C2E.asnumpy()[:, C2E_neighbor_idx], 1:] + 1.0), axis=1 ) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py index 9fab169ef1..b08ff90f4f 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py @@ -21,6 +21,13 @@ from . import util + +class boden(gtx_common.DimensionIndex): ... + + +class K(gtx_common.DimensionIndex, kind=gtx_common.DimensionKind.VERTICAL): ... + + N = 10 @@ -312,10 +319,8 @@ def _make_horizontal_promoter_sdfg( sdfg = dace.SDFG(util.unique_name("serial_map_promoter_tester")) state = sdfg.add_state(is_start_block=True) - h_idx = gtx_dace_lowering.get_map_variable(gtx_common.Dimension("boden")) - v_idx = gtx_dace_lowering.get_map_variable( - gtx_common.Dimension("K", gtx_common.DimensionKind.VERTICAL) - ) + h_idx = gtx_dace_lowering.get_map_variable(boden) + v_idx = gtx_dace_lowering.get_map_variable(K) if d1_map_is_vertical: d1_shape = (10,) diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index 84fae1145e..a6aad1cb56 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -18,6 +18,7 @@ import gt4py.next.common as common from gt4py.next.common import ( Dimension, + DimensionIndex, DimensionKind, Domain, Infinity, @@ -29,15 +30,58 @@ unit_range, ) -C2E = Dimension("C2E", kind=DimensionKind.LOCAL) -V2E = Dimension("V2E", kind=DimensionKind.LOCAL) -E2V = Dimension("E2V", kind=DimensionKind.LOCAL) -E2C = Dimension("E2C", kind=DimensionKind.LOCAL) -E2C2V = Dimension("E2C2V", kind=DimensionKind.LOCAL) -ECDim = Dimension("ECDim") -IDim = Dimension("IDim") -JDim = Dimension("JDim") -KDim = Dimension("KDim", kind=DimensionKind.VERTICAL) + +class X(DimensionIndex): ... + + +class Y(DimensionIndex): ... + + +class Z(DimensionIndex): ... + + +class Foo(DimensionIndex): ... + + +class J(DimensionIndex): ... + + +class K(DimensionIndex): ... + + +class I(common.DimensionIndex): ... + + +class I_half(common.DimensionIndex): ... + + +class C2E(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class V2E(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class E2C2V(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class ECDim(DimensionIndex): ... + + +class IDim(DimensionIndex): ... + + +class JDim(DimensionIndex): ... + + +class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... + + IHalfDim = common.flip_staggered(IDim) @@ -394,7 +438,7 @@ def test_domain_dims_ranges_length_mismatch(): ValueError, match=r"Number of provided dimensions \(\d+\) does not match number of provided ranges \(\d+\)", ): - dims = [Dimension("X"), Dimension("Y"), Dimension("Z")] + dims = [X, Y, Z] ranges = [UnitRange(0, 1), UnitRange(0, 1)] Domain(dims=dims, ranges=ranges) @@ -435,21 +479,21 @@ def test_domain_slice_at(): def test_domain_dim_index(): - dims = [Dimension("X"), Dimension("Y"), Dimension("Z")] + dims = [X, Y, Z] ranges = [UnitRange(0, 1), UnitRange(0, 1), UnitRange(0, 1)] domain = Domain(dims=dims, ranges=ranges) - domain.dim_index(Dimension("Y")) == 1 + domain.dim_index(Y) == 1 - domain.dim_index(Dimension("Foo")) == None + domain.dim_index(Foo) == None def test_domain_pop(): - dims = [Dimension("X"), Dimension("Y"), Dimension("Z")] + dims = [X, Y, Z] ranges = [UnitRange(0, 1), UnitRange(0, 1), UnitRange(0, 1)] domain = Domain(dims=dims, ranges=ranges) - domain.pop(Dimension("X")) == Domain(dims=dims[1:], ranges=ranges[1:]) + domain.pop(X) == Domain(dims=dims[1:], ranges=ranges[1:]) domain.pop(0) == Domain(dims=dims[1:], ranges=ranges[1:]) @@ -462,92 +506,92 @@ def test_domain_pop(): # Valid index and named ranges ( 0, - [NamedRange(Dimension("X"), UnitRange(100, 110))], + [NamedRange(X, UnitRange(100, 110))], Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), Domain( - NamedRange(Dimension("X"), UnitRange(100, 110)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(X, UnitRange(100, 110)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), ), ( 1, - [NamedRange(Dimension("X"), UnitRange(100, 110))], + [NamedRange(X, UnitRange(100, 110))], Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("X"), UnitRange(100, 110)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(X, UnitRange(100, 110)), + NamedRange(K, UnitRange(0, 10)), ), ), ( -1, - [NamedRange(Dimension("X"), UnitRange(100, 110))], + [NamedRange(X, UnitRange(100, 110))], Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("X"), UnitRange(100, 110)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(X, UnitRange(100, 110)), ), ), ( - Dimension("J"), + J, [ - NamedRange(Dimension("X"), UnitRange(100, 110)), - NamedRange(Dimension("Z"), UnitRange(100, 110)), + NamedRange(X, UnitRange(100, 110)), + NamedRange(Z, UnitRange(100, 110)), ], Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("X"), UnitRange(100, 110)), - NamedRange(Dimension("Z"), UnitRange(100, 110)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(X, UnitRange(100, 110)), + NamedRange(Z, UnitRange(100, 110)), + NamedRange(K, UnitRange(0, 10)), ), ), # Invalid indices ( 3, - [NamedRange(Dimension("X"), UnitRange(100, 110))], + [NamedRange(X, UnitRange(100, 110))], Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), IndexError, ), ( -4, - [NamedRange(Dimension("X"), UnitRange(100, 110))], + [NamedRange(X, UnitRange(100, 110))], Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), IndexError, ), ( - Dimension("Foo"), - [NamedRange(Dimension("X"), UnitRange(100, 110))], + Foo, + [NamedRange(X, UnitRange(100, 110))], Domain( - NamedRange(Dimension("I"), UnitRange(0, 10)), - NamedRange(Dimension("J"), UnitRange(0, 10)), - NamedRange(Dimension("K"), UnitRange(0, 10)), + NamedRange(I, UnitRange(0, 10)), + NamedRange(J, UnitRange(0, 10)), + NamedRange(K, UnitRange(0, 10)), ), ValueError, ), @@ -667,7 +711,6 @@ def test_hashes(self): class TestCartesianConnectivity: def test_for_translation(self): offset = 5 - I = common.Dimension("I") result = common.CartesianConnectivity.for_translation(I, offset) assert isinstance(result, common.CartesianConnectivity) @@ -676,8 +719,6 @@ def test_for_translation(self): assert result.offset == offset def test_for_relocation(self): - I = common.Dimension("I") - I_half = common.Dimension("I_half") result = common.CartesianConnectivity.for_relocation(I, I_half) assert isinstance(result, common.CartesianConnectivity) diff --git a/tests/next_tests/unit_tests/test_constructors.py b/tests/next_tests/unit_tests/test_constructors.py index 0d6bbdeb2b..3bfdacbf63 100644 --- a/tests/next_tests/unit_tests/test_constructors.py +++ b/tests/next_tests/unit_tests/test_constructors.py @@ -22,9 +22,14 @@ ) -I = gtx.Dimension("I") -J = gtx.Dimension("J") -K = gtx.Dimension("K") +class I(gtx.DimensionIndex): ... + + +class J(gtx.DimensionIndex): ... + + +class K(gtx.DimensionIndex): ... + sizes = {I: 10, J: 10, K: 10} @@ -254,7 +259,7 @@ def test_array_namespace_allocator_with_device(self): def test_array_namespace_allocator_aligned_index_warns(self): """allocator=numpy with aligned_index → warns and ignores aligned_index.""" - aligned_index = [common.NamedIndex(I, 0)] + aligned_index = [I(0)] with pytest.warns(UserWarning, match="aligned_index"): fc = constructors.FieldConstructor(allocator=np, aligned_index=aligned_index) field = fc.zeros(self._domain) @@ -272,7 +277,7 @@ def test_field_buffer_allocator(self): def test_field_buffer_allocator_with_aligned_index(self): """allocator=FieldBufferAllocator + aligned_index → accepted without error.""" allocator = next_allocators.StandardCPUFieldBufferAllocator() - aligned_index = [common.NamedIndex(I, 0)] + aligned_index = [I(0)] fc = constructors.FieldConstructor(allocator=allocator, aligned_index=aligned_index) field = fc.zeros(self._domain) assert isinstance(field.ndarray, np.ndarray) 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 6f54a5d930..f5257506fb 100644 --- a/tests/next_tests/unit_tests/test_custom_layout_allocators.py +++ b/tests/next_tests/unit_tests/test_custom_layout_allocators.py @@ -17,6 +17,30 @@ import gt4py.storage.allocators as core_allocators +class D0(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class D1(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class D2(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class D0_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + + +class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class D2_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + + +class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class D1_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + + class DummyAllocator(next_allocators.FieldBufferAllocatorProtocol): __gt_device_type__ = core_defs.DeviceType.CPU @@ -25,7 +49,7 @@ def __gt_allocate__( domain: common.Domain, dtype: core_defs.DType[core_defs.ScalarT], device_id: int = 0, - aligned_index: Optional[Sequence[common.NamedIndex]] = None, + aligned_index: Optional[Sequence[common.DimensionIndex]] = None, ) -> core_allocators.TensorBuffer[core_defs.DeviceTypeT, core_defs.ScalarT]: pass @@ -98,35 +122,39 @@ def test_horizontal_first_layout_mapper(): # Test with only horizontal dimensions dims = [ - common.Dimension("D0", common.DimensionKind.HORIZONTAL), - common.Dimension("D1", common.DimensionKind.HORIZONTAL), - common.Dimension("D2", common.DimensionKind.HORIZONTAL), + D0, + D1, + D2, ] expected_layout_map = core_allocators.BufferLayoutMap((2, 1, 0)) assert horizontal_first_layout_mapper(dims) == expected_layout_map # Test with no horizontal dimensions dims = [ - common.Dimension("D0", common.DimensionKind.VERTICAL), - common.Dimension("D1", common.DimensionKind.LOCAL), - common.Dimension("D2", common.DimensionKind.VERTICAL), + D0_vertical, + D1_local, + D2_vertical, ] expected_layout_map = core_allocators.BufferLayoutMap((2, 0, 1)) assert horizontal_first_layout_mapper(dims) == expected_layout_map # Test with a mix of dimensions dims = [ - common.Dimension("D2", common.DimensionKind.LOCAL), - common.Dimension("D0", common.DimensionKind.HORIZONTAL), - common.Dimension("D1", common.DimensionKind.VERTICAL), + D2_local, + D0, + D1_vertical, ] expected_layout_map = core_allocators.BufferLayoutMap((0, 2, 1)) assert horizontal_first_layout_mapper(dims) == expected_layout_map -Cell = common.Dimension("Cell", common.DimensionKind.HORIZONTAL) -Edge = common.Dimension("Edge", common.DimensionKind.HORIZONTAL) -K = common.Dimension("K", common.DimensionKind.VERTICAL) +class Cell(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... class TestBaseFieldBufferAllocatorAlignedIndex: @@ -151,7 +179,7 @@ def test_aligned_index_zero_origin(self): domain = common.Domain( dims=(Cell, K), ranges=(common.UnitRange(0, 10), common.UnitRange(0, 5)) ) - aligned = [common.NamedIndex(Cell, 3)] + aligned = [Cell(3)] result = allocator.__gt_allocate__(domain, core_defs.dtype(float), aligned_index=aligned) assert result.shape == (10, 5) assert result.aligned_index == (3, 0) @@ -162,7 +190,7 @@ def test_aligned_index_nonzero_origin_cell(self): domain = common.Domain( dims=(Cell, K), ranges=(common.UnitRange(100, 110), common.UnitRange(0, 5)) ) - aligned = [common.NamedIndex(Cell, 103)] + aligned = [Cell(103)] result = allocator.__gt_allocate__(domain, core_defs.dtype(float), aligned_index=aligned) assert result.shape == (10, 5) assert result.aligned_index == (3, 0) @@ -173,7 +201,7 @@ def test_aligned_index_nonzero_origin_edge(self): domain = common.Domain( dims=(Edge, K), ranges=(common.UnitRange(200, 212), common.UnitRange(0, 5)) ) - aligned = [common.NamedIndex(Edge, 207)] + aligned = [Edge(207)] result = allocator.__gt_allocate__(domain, core_defs.dtype(float), aligned_index=aligned) assert result.shape == (12, 5) assert result.aligned_index == (7, 0) @@ -184,7 +212,7 @@ def test_aligned_index_nonzero_origin_all_dims(self): domain = common.Domain( dims=(Cell, K), ranges=(common.UnitRange(50, 60), common.UnitRange(10, 15)) ) - aligned = [common.NamedIndex(Cell, 53), common.NamedIndex(K, 12)] + aligned = [Cell(53), K(12)] result = allocator.__gt_allocate__(domain, core_defs.dtype(float), aligned_index=aligned) assert result.shape == (10, 5) assert result.aligned_index == (3, 2) @@ -195,7 +223,7 @@ def test_aligned_index_at_domain_start(self): domain = common.Domain( dims=(Cell, K), ranges=(common.UnitRange(100, 110), common.UnitRange(20, 25)) ) - aligned = [common.NamedIndex(Cell, 100), common.NamedIndex(K, 20)] + aligned = [Cell(100), K(20)] result = allocator.__gt_allocate__(domain, core_defs.dtype(float), aligned_index=aligned) assert result.shape == (10, 5) assert result.aligned_index == (0, 0) @@ -204,7 +232,7 @@ def test_aligned_index_shared_between_cell_and_edge_fields(self): """Same aligned_index with both Cell and Edge can be used to allocate both a Cell-field and an Edge-field; the irrelevant dimension is ignored.""" allocator = self._make_allocator() - aligned = [common.NamedIndex(Cell, 103), common.NamedIndex(Edge, 207)] + aligned = [Cell(103), Edge(207)] cell_domain = common.Domain( dims=(Cell, K), ranges=(common.UnitRange(100, 110), common.UnitRange(0, 5)) @@ -236,7 +264,7 @@ def test_aligned_index_outside_domain_raises(self): allocator.__gt_allocate__( domain, core_defs.dtype(float), - aligned_index=[common.NamedIndex(Cell, 50)], + aligned_index=[Cell(50)], ) # After domain end @@ -244,7 +272,7 @@ def test_aligned_index_outside_domain_raises(self): allocator.__gt_allocate__( domain, core_defs.dtype(float), - aligned_index=[common.NamedIndex(Cell, 120)], + aligned_index=[Cell(120)], ) # Exactly at domain end (exclusive upper bound) @@ -252,5 +280,5 @@ def test_aligned_index_outside_domain_raises(self): allocator.__gt_allocate__( domain, core_defs.dtype(float), - aligned_index=[common.NamedIndex(Cell, 110)], + aligned_index=[Cell(110)], ) diff --git a/tests/next_tests/unit_tests/test_field_utils.py b/tests/next_tests/unit_tests/test_field_utils.py index ffe7feedef..a26ffa2992 100644 --- a/tests/next_tests/unit_tests/test_field_utils.py +++ b/tests/next_tests/unit_tests/test_field_utils.py @@ -12,6 +12,9 @@ from gt4py.next import common, constructors, field_utils +class X(common.DimensionIndex): ... + + @pytest.mark.parametrize( "device_type", [ @@ -23,7 +26,7 @@ def test_verify_device_field_type(nd_array_implementation_and_device_type, device_type): nd_array_implementation, compatible_device_type = nd_array_implementation_and_device_type - testee = constructors.as_field([common.Dimension("X")], nd_array_implementation.asarray([42.0])) + testee = constructors.as_field([X], nd_array_implementation.asarray([42.0])) is_correct_device = compatible_device_type == device_type assert field_utils.verify_device_field_type(testee, device_type) == is_correct_device diff --git a/tests/next_tests/unit_tests/test_utils.py b/tests/next_tests/unit_tests/test_utils.py index 6068d06f3d..6ad27a6f59 100644 --- a/tests/next_tests/unit_tests/test_utils.py +++ b/tests/next_tests/unit_tests/test_utils.py @@ -13,11 +13,17 @@ import pytest from gt4py.eve import concepts, datamodels -from gt4py.next import fingerprinting, utils +from gt4py.next import common, fingerprinting, utils from eve_tests import definitions +class I(common.DimensionIndex): ... + + +class J(common.DimensionIndex): ... + + @dataclasses.dataclass class _DataclassModel: value: int @@ -427,7 +433,7 @@ def test_dicts_with_unorderable_keys_are_order_independent(self): # `Dimension`s (and other dataclasses without `__lt__`) occur as dict keys # e.g. in user closure variables. - i, j = common.Dimension("I"), common.Dimension("J") + i, j = I, J assert fingerprinting.strict_fingerprinter( {i: 1, j: 2} ) == fingerprinting.strict_fingerprinter({j: 2, i: 1}) 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 14a90409bd..6ee13a358c 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 @@ -12,13 +12,30 @@ from gt4py.next import ( Dimension, + DimensionIndex, DimensionKind, ) from gt4py.next.type_system import type_info, type_specifications as ts from gt4py.next.ffront import type_specifications as ts_ffront from gt4py.next.iterator.type_system import type_specifications as ts_it -TDim = Dimension("TDim") # Meaningless dimension, used for tests. + +class IDim(DimensionIndex): ... + + +class JDim(DimensionIndex): ... + + +class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... + + +class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class TDim(DimensionIndex): ... def type_info_cases() -> list[tuple[Optional[ts.TypeSpec], dict]]: @@ -61,14 +78,10 @@ def callable_type_info_cases(): if not isinstance(symbol_type, ts.CallableType) ] - IDim = Dimension("I") - JDim = Dimension("J") - KDim = Dimension("K", kind=DimensionKind.VERTICAL) - bool_type = ts.ScalarType(kind=ts.ScalarKind.BOOL) float_type = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) int_type = ts.ScalarType(kind=ts.ScalarKind.INT64) - field_type = ts.FieldType(dims=[Dimension("I")], dtype=float_type) + field_type = ts.FieldType(dims=[IDim], dtype=float_type) tuple_type = ts.TupleType(types=[bool_type, field_type]) nullary_func_type = ts.FunctionType( pos_only_args=[], pos_or_kw_args={}, kw_only_args={}, returns=ts.VoidType() @@ -261,7 +274,7 @@ def callable_type_info_cases(): [ts.TupleType(types=[float_type, field_type])], {}, [ - r"Expected 1st argument to be of type 'tuple\[bool, Field\[\[I\], float64\]\]', got 'tuple\[float64, Field\[\[I\], float64\]\]'" + r"Expected 1st argument to be of type 'tuple\[bool, Field\[\[IDim\], float64\]\]', got 'tuple\[float64, Field\[\[IDim\], float64\]\]'" ], ts.VoidType(), ), @@ -270,7 +283,7 @@ def callable_type_info_cases(): [int_type], {}, [ - r"Expected 1st argument to be of type 'tuple\[bool, Field\[\[I\], float64\]\]', got 'int64'" + r"Expected 1st argument to be of type 'tuple\[bool, Field\[\[IDim\], float64\]\]', got 'int64'" ], ts.VoidType(), ), @@ -292,8 +305,8 @@ def callable_type_info_cases(): ], {}, [ - r"Expected argument 'a' to be of type 'Field\[\[K\], int64\]', got 'Field\[\[K\], float64\]'", - r"Expected argument 'b' to be of type 'Field\[\[K\], int64\]', got 'Field\[\[K\], float64\]'", + r"Expected argument 'a' to be of type 'Field\[\[KDim\], int64\]', got 'Field\[\[KDim\], float64\]'", + r"Expected argument 'b' to be of type 'Field\[\[KDim\], int64\]', got 'Field\[\[KDim\], float64\]'", ], ts.FieldType(dims=[KDim], dtype=float_type), ), @@ -343,8 +356,8 @@ def callable_type_info_cases(): [ts.TupleType(types=[ts.FieldType(dims=[IDim, JDim, KDim], dtype=int_type)])], {}, [ - r"Expected argument 'a' to be of type 'tuple\[Field\[\[I, J, K\], int64\], " - r"Field\[\[\.\.\.\], int64\]\]', got 'tuple\[Field\[\[I, J, K\], int64\]\]'." + r"Expected argument 'a' to be of type 'tuple\[Field\[\[IDim, JDim, KDim\], int64\], " + r"Field\[\[\.\.\.\], int64\]\]', got 'tuple\[Field\[\[IDim, JDim, KDim\], int64\]\]'." ], ts.FieldType(dims=[IDim, JDim, KDim], dtype=float_type), ), @@ -412,16 +425,14 @@ def test_return_type( [ (ts.ScalarType(kind=ts.ScalarKind.INT64), False), ( - ts.FieldType(dims=[Dimension("I")], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)), + ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)), False, ), ( ts.TupleType( types=[ ts.ScalarType(kind=ts.ScalarKind.INT64), - ts.FieldType( - dims=[Dimension("I")], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64) - ), + ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)), ] ), False, @@ -430,9 +441,7 @@ def test_return_type( ts.NamedCollectionType( types=[ ts.ScalarType(kind=ts.ScalarKind.INT64), - ts.FieldType( - dims=[Dimension("I")], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64) - ), + ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)), ], keys=["a", "b"], original_python_type="some.module:SomeClass", @@ -447,7 +456,7 @@ def test_return_type( types=[ ts.ScalarType(kind=ts.ScalarKind.INT64), ts.FieldType( - dims=[Dimension("I")], + dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64), ), ], @@ -467,8 +476,6 @@ def test_needs_value_extraction(type_spec: ts.TypeSpec, expected: bool): def test_promote_lists(): float64 = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) int32 = ts.ScalarType(kind=ts.ScalarKind.INT32) - V2EDim = Dimension("V2E", kind=DimensionKind.LOCAL) - C2EDim = Dimension("C2E", kind=DimensionKind.LOCAL) const_list = ts.ListType(element_type=float64, offset_type=None) v2e_list = ts.ListType(element_type=float64, offset_type=V2EDim) diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py index 08783b6efd..7ba7d09b5b 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py @@ -30,8 +30,11 @@ def dtype(self) -> np.dtype: return np.dtype(np.int32) -IDim = gtx.Dimension("IDim") -JDim = gtx.Dimension("JDim") +class IDim(gtx.DimensionIndex): ... + + +class JDim(gtx.DimensionIndex): ... + # -- PEP 695 type aliases -- type IFloatFieldAlias = gtx.Field[gtx.Dims[IDim], float] @@ -100,7 +103,7 @@ def _make_type_string_for_container(cls: type) -> str: ( gtx.Field[[IDim, JDim], float], ts.FieldType( - dims=[gtx.Dimension("IDim"), gtx.Dimension("JDim")], + dims=[IDim, JDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64), ), ), From 05f14f1f434fa24f17c4e68c7065eb203e8c6243 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 12:56:50 +0200 Subject: [PATCH 05/20] docs[next]: drop the implementation plan from the PR It was committed by accident, via a broad `git add` of `docs/`. It is a working artifact for planning this stack -- with four rounds of review notes folded in -- not documentation for the repository, and ADR 0028 already records the decisions that need to live here. It stays available outside the PR. --- ...ectivities-as-types-implementation-plan.md | 1033 ----------------- 1 file changed, 1033 deletions(-) delete mode 100644 docs/development/next/connectivities-as-types-implementation-plan.md diff --git a/docs/development/next/connectivities-as-types-implementation-plan.md b/docs/development/next/connectivities-as-types-implementation-plan.md deleted file mode 100644 index 52bdbbc0c1..0000000000 --- a/docs/development/next/connectivities-as-types-implementation-plan.md +++ /dev/null @@ -1,1033 +0,0 @@ -# Connectivities as types — implementation plan - -**Status**: **APPROVED** at revision 6 (adversarial review rounds 1–4; round 4 verdict APPROVED) -**Target**: an 8-PR stack on `main`, *alternative to* GridTools/gt4py#2844 -**Proposal**: `egparedes/connectivities-as-types` in GridTools/gt4py_knowledge (PR #32) -**Baseline tree**: `b3c53fa7e` (v1.2.2) - -## 0. Scope and relation to #2844 - -The proposal and #2844 agree on *what a dimension is* (a class, its indices its -instances) and disagree on *what identity a dimension has*. #2844 chose -`(tag, kind)` value equality with an interning registry; the proposal chose -nominal type identity with the tag being the qualified Python name. That single -disagreement propagates into five mechanisms, so the two cannot both land. - -This stack **re-cuts** #2844: it keeps the machinery independent of identity, -drops the machinery that exists only to support value identity, and then builds -the connectivity layer on top. **#2844 is closed, not merged** — that is what -makes this an alternative. Consequences: the new dimension ADR is **0028** (the -ADR directory on `main` ends at **0027**; 0028 exists only inside unmerged -#2844), and there is nothing to supersede. - -### Taken from #2844, unchanged in substance - -| Piece | Where in #2844 | -| --- | --- | -| `DimensionMeta` metaclass; `I + 1`, `I > 5`, `repr` living on it | `common.py` | -| `DimensionIndex` base: `__slots__ = ("value",)`, `kind` class keyword, `.dim` property | `common.py` | -| `type Dimension = type[DimensionIndex]` as a PEP 695 alias (so `Dimension("I")` raises rather than silently evaluating to `str`) | `common.py` | -| Metaclass `.value` property raising `AttributeError` that points at `.tag` | `common.py` | -| Deletion of `common.NamedIndex` (`.dim` / `.value` move onto the index instance) | `common.py` + ~40 call sites | -| Deletion of the dimension half of `mypy_plugin.py` (`_DimA`..`_AnyDim`); only the mixed-precision hooks remain | `type_system/mypy_plugin.py` | -| The mechanical migration of every declaration, incl. docs, workshop notebooks and `examples/` (which `test_examples` executes) | 131 files: 52 `src/`, 66 `tests/`, 13 docs | -| `xtyping.resolve_annotation` usage at `fbuiltins._type_conversion_helper` (already on `main` via #2841) | — | - -### Dropped from #2844 - -| Piece | Why | -| --- | --- | -| `_DIMENSION_REGISTRY` interning | identity is the type; nothing to intern | -| `copyreg.pickle(DimensionMeta, _reduce_dimension)` — the **blanket** registration on all dimensions | verified: a module-level dimension class pickles by reference with no help (`pickle.loads(pickle.dumps(KDim)) is KDim`). **But a narrow `copyreg` on `StaggeredMeta` is still required** — see §1.5 | -| The `DimensionMeta`-vs-`DimensionMeta` branch of `__eq__` / `__ne__` | becomes `is`. **The `IntegralScalar` overloads (`I == 5` → `Domain`) are kept**, and therefore so is an explicit `__hash__` — see §1.0 | -| `DimensionIndex.__eq__` comparing `type(self) == type(other)` | becomes `type(self) is type(other)` | -| `common.dimension(tag, kind)` factory | replaced by `common.resolve(tag)`, which imports | -| `fingerprinting.py` `DimensionMeta` deconstructor keyed on `(tag, kind)` | under type identity a dimension *is* fingerprinted by qualified name, so the generic `type` deconstructor is correct — **for the lenient variant only**. The STRICT variant rejects `Staggered[KDim]`, which is not importable under its qualified name. Both in-tree fingerprinters are lenient (`ffront/stages.py:62`, `iterator/ir.py:26`), and `eve_utils.content_hash` (`compiled_program.py:420`) is pickle-based and so needs §1.5's `copyreg`. Record the STRICT caveat in the ADR | -| ADR 0028 as drafted in #2844 | never lands; this stack writes its own 0028 | - -### Changed relative to #2844 - -| Piece | #2844 | This stack | -| --- | --- | --- | -| `tag` default | `cls.__name__`, settable in the class body | `f"{cls.__module__}.{cls.__qualname__}"`, a metaclass property; a class-body `tag = ...` is a `TypeError` | -| Rebuilding a dimension from a tag | `dimension(tag, kind)` (registry) | `resolve(tag)` (`import_module` + `qualname` walk), memoized | -| Declaration site requirement | none | module level, or unpicklable; `` heuristic in `__init_subclass__` | -| Backend name mangling | `tag` used directly | `codegen_name(tag)` + inverse, at ~19 enumerated sites in two name spaces (§1.3(b), (c)) | -| Staggered dimensions | `_Staggered` prefix through the interning factory | `Staggered[D]`, a real parametrized type — **required in PR 2**, not optional (§1.5) | - -### Superseded - -- **#2845** (`test[next]: adopt class-style dimension declarations`) is subsumed - by PR 2: because `dimension()` is not user-facing, the minimal - `I = gtx.dimension("I")` form does not exist and every declaration takes class - form immediately. #2845's pyright coverage is folded in. -- The `FieldOffset`-as-frontend-identifier part of **ADR 0019**. -- **ADR 0026**'s `_Staggered` name prefix (PR 2). - -## 1. Design questions closed before implementation - -Everything in this section was verified by running it, not by reading. Probe -files are named; they become committed test material in the PR that needs them. - -### 1.0 Metaclass mechanics that are easy to get wrong - -**`__hash__` must be declared explicitly.** Python sets `__hash__ = None` on any -class body that defines `__eq__` without `__hash__` — metaclasses included. Since -the `I == 5` → `Domain` overload keeps `__eq__` on `DimensionMeta`, dropping -#2844's `__hash__` makes every dimension class *unhashable*: - -``` ->>> class M(type): -... def __eq__(cls, o): return True ->>> M.__hash__ is None -True ->>> class C(metaclass=M): pass ->>> hash(C) -TypeError: unhashable type: 'M' -``` - -That would break `domain({I: 2})` (`common.py:672-690`), -`Counter[common.Dimension]` (`embedded/nd_array_field.py:314`), -`dict[Dimension, SymbolicRange]` (`iterator/ir_utils/domain_utils.py:136,152`), -`seen: dict[Dimension, Dimension]` (`common.py:1351`), and eve's validator -memoization on annotation objects (`eve/type_validation.py:599`) — so -`ts.DimensionType` would fail at *import*. **Fix: `__hash__ = type.__hash__` -explicitly on `DimensionMeta`,** and likewise on `ConnectivityMeta` if it ever -defines `__eq__`. - -**A metaclass `__getitem__` shadows `__class_getitem__`.** `ConnectivityMeta` -needs `__getitem__` for `V2E[1]` (the single-neighbor shift handle that -`FieldOffset.__getitem__` provides today), but metaclass lookup takes precedence -over `Generic.__class_getitem__`, so a naive implementation makes -`NeighborConnectivity[V, E]` in a bases list fail with -`TypeError: tuple expected at most 1 argument, got 3`. - -**Fix, verified clean under `mypy --strict` and pyright 1.1.414 on Python 3.12** -(`/tmp/probe_meta_getitem3.py`): dispatch on the argument type, delegating -non-`int` subscription back to `cls.__class_getitem__`: - -```python -class ConnectivityMeta(type): - __hash__ = type.__hash__ - @overload - def __getitem__(cls, item: int) -> Connectivity: ... - @overload - def __getitem__(cls, item: Any) -> Any: ... - def __getitem__(cls, item: Any) -> Any: - # `numbers.Integral`, not `int`: `V2E[np.int32(1)]` must not fall through - # to the type-parameter branch (it raises `TypeError: V2E is not a - # generic class` there). `bool` is excluded so `V2E[True]` is an error - # rather than silently neighbor 1. - if isinstance(item, numbers.Integral) and not isinstance(item, bool): - return _bound_single_neighbor(cls, int(item)) - # type-parameter subscription, e.g. `NeighborConnectivity[V, E]` - return cast(Any, cls).__class_getitem__(item) -``` - -`cast(Any, cls)`, not `super()` — `__class_getitem__` is on the class, not on the -metaclass MRO; `super().__class_getitem__` raises `AttributeError`. With the cast -both checkers report zero errors and all uses work at runtime -(`NC[V, E]`, `class V2E(NC[V, E])`, `V2E[1]`, and `V2E.Local` as an annotation). -pyright accepts `NC[V, E]` in a **bases list**; in a *value* position it types it -`Any`, which is why the overloads above matter — without them `V2E[1]` is also -`Any` and the shift handle is untyped. - -### 1.1 `NeighborConnectivity` is **not** a `Connectivity` (proposal Open Q6) - -`common.Connectivity` is `Field[DimsT, IntegralScalar]` — a **data** protocol -(`common.py:990`; `ndarray`, `asnumpy`, `domain` are all on it). A declaration -class holds no data. - -**Resolution.** Two distinct things, distinct hierarchies: - -- `NeighborConnectivity` — a **declaration**. Not a `Connectivity`. It produces a - `NeighborConnectivityType` via `__gt_type__()`, is the provider key, and is the - handle written in DSL code (`a(V2E)`). -- `NeighborTable` / `NdArrayConnectivityField` — the **data**, unchanged, still - `Connectivity` implementations. - -This is the shape `FieldOffset` already has: it is *not* a `Connectivity` either, -and `premap` special-cases it at `nd_array_field.py:317-320`. So `a(V2E)` -continues to work by widening the same union — `Field.premap` and -`Field.__call__` are typed `Connectivity | fbuiltins.FieldOffset` -(`common.py:785, 791-794`) and become `Connectivity | type[NeighborConnectivity]` -in PR 4. `V2E` has **no instances**: `ConnectivityMeta.__call__` raises -`TypeError("… is a connectivity declaration and cannot be instantiated; bind a -table through the offset provider")`. - -Consequence: the proposal's sketch line -`class NeighborConnectivity(Connectivity[MultiDimensionIndex[Origin, Local], Codomain], ...)` -is **wrong and dropped**. `MultiDimensionIndex` remains the *domain index type* of -the `NeighborTable` (PR 8). **The knowledge-repo note needs this correction.** - -### 1.2 How `Local` reaches the base (proposal Open Q2) - -`requires-python = '>=3.12'`, so a PEP 696 default type parameter (3.13) is not -available. Resolution: **metaclass discovery**, base carrying a `ClassVar` -annotation, subclass declaring the nested class explicitly: - -```python -class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( - metaclass=ConnectivityMeta -): - Local: ClassVar[type[LocalDimensionIndex]] # annotation only, never assigned - -class V2E(NeighborConnectivity[V, E], max_neighbors=6): - class Local(LocalDimensionIndex): ... # explicit, required -``` - -Verified under `mypy --strict --python-version 3.12` and `pyright --pythonversion 3.12` -(`/tmp/probe_local.py`): - -| Variant | base declares | mypy | pyright | -| --- | --- | --- | --- | -| 1 | `Local: ClassVar[type[LocalDimensionIndex]]` | clean | clean | -| 2 | nothing | clean | clean | -| 3 | a real nested `class Local(LocalDimensionIndex)` | clean | **`reportIncompatibleVariableOverride`** | - -In all three the intended negative case (`Field[V, A.Local]` vs -`Field[V, B.Local]`) is correctly an error. Variant 3 is rejected. - -**Stated precisely — what variant 1 does and does not buy.** It does *not* make -`conn.Local` usable as a **type annotation** when `conn` is a generic -`type[NeighborConnectivity]`: both checkers reject that (mypy `name-defined`, -pyright `reportInvalidTypeForm`), and `T.Local` on a `TypeVar` is rejected too. -What variant 1 buys over variant 2 is only **value-level** access — -`reveal_type(conn.Local)` is `type[LocalDimensionIndex]` instead of an attribute -error — which is what library code in `common`, the backends and -`type_synthesizer` actually needs. Variant 1 is chosen for that, not for generic -annotations. Generic library code that must *name* a local dimension in a -signature uses `type[LocalDimensionIndex]`. - -This extends the proposal's probe P2: a **generated** `Local` is unusable as an -annotation, but a base `ClassVar` *annotation* plus an explicitly declared nested -class is fine. - -### 1.3 The IR keeps string tags; `resolve()` and `codegen_name()` are both required - -`AxisLiteral.value: str` stays (making it carry the class is a separate IR -change, deferred past this stack). It now holds the **qualified** tag, and that -has two consequences the first draft of this plan underestimated. - -**(a) `resolve(tag)` at every rebuild site**, memoized — `inference.py:464` calls -it once per `AxisLiteral` on the type-inference hot path: - -| Site | Purpose | -| --- | --- | -| `iterator/ir_utils/domain_utils.py` | `AxisLiteral` → `Dimension` | -| `iterator/ir_utils/misc.py` | `AxisLiteral` → `Dimension` | -| `iterator/type_system/inference.py:464` | `AxisLiteral` → `ts.DimensionType` | -| `codegens/gtfn/itir_to_gtfn_ir.py` (×2) | staggered-name sniffing → replaced in PR 2 by `Staggered[D]` | -| `dace/lowering/gtir_to_sdfg_lambda.py:1155` | synthesizes the local dim from the *offset* tag: `Dimension(offset, LOCAL)`. **Must be fixed in PR 2, not deferred** — see below | -| `dace/sdfg_args.py` | axis name → `Dimension` | -| `runners/roundtrip.py` | emits `gtx.Dimension(...)` as *source text* → becomes an import | -| ~~`common.flip_staggered` (×2)~~ | **not** a `resolve()` site: `Staggered[D]` replaces it with an interning subscript, §1.5 | - -`resolve` on a nested qualname was verified to work and round-trip -(`resolve("mymod.V2E.Local") is mymod.V2E.Local`), which matters because PR 4 keys -the provider on `V2E.Local.tag`. **One hazard to settle in PR 2**: a purely dotted -tag does not record *where* the module path ends and the qualname begins, so -`resolve` must try the longest importable prefix and walk the rest — O(depth) -import attempts, and in principle ambiguous if a module path and a class-attribute -chain collide. `pickle` avoids this by storing module and qualname *separately*. -Options: keep the pure dotted form (what the proposal asks for, ambiguity -tolerated and memoized away) or use an explicit separator such as -`"module:qualname"`. **Recommendation: keep the dotted form** — it is what makes -the tag "also a valid tag string for the IR", the collision requires a module and -an attribute chain to have the same spelling, and `resolve` can prefer the -*longest* importable prefix so a real module always wins. Record the residual in -the ADR. - -**(b) `codegen_name(tag)` — dots are illegal in every generated identifier.** -`eve`'s `SymbolName`/`SymbolRef` are constrained by -`_SYMBOL_NAME_RE = ^[a-zA-Z_]\w*$` (`eve/concepts.py:23,26,32`), so a qualified -tag reaching `Sym(id=...)` is a *validation error*, not a cosmetic problem. The -first draft mentioned mangling only in the abstract and put the roundtrip change -in a later PR; both were wrong. All of these are **PR 2**: - -| Site | What breaks without mangling | -| --- | --- | -| `codegens/gtfn/itir_to_gtfn_ir.py:170-195` | `TagDefinition(name=Sym(id=dim.value))` → `SymbolName` validation error | -| `codegens/gtfn/gtfn_module.py:97, 130-136` | `generated::{dim.value}_t`, plus `name.lower()` | -| `otf/binding/nanobind.py:197, 211` | C++ identifiers | -| `dace/lowering/gtir_to_sdfg_utils.py` `get_map_variable` | `i_{dim.value}_gtx_{kind}` → invalid DaCe symbol | -| `dace/sdfg_args.py:80` `_field_symbol` | invalid DaCe symbol | -| `dace/lowering/gtir_python_codegen.py:137-138` | `visit_AxisLiteral` returns the raw value | -| `runners/roundtrip.py:64, 177` | `AxisLiteral = as_fmt("{value}")`, and `{o.value} = gtx.Dimension(...)` emits `a.b.I = ...` → `SyntaxError` | - -**(d) The mangling scheme, corrected.** Earlier drafts said "injective (escape -existing `__` before replacing `.`)", i.e. `_ -> __` then `. -> _`. **That is not -injective**: `.` becomes a single `_`, so `".."` and `"_"` both map to `"__"`. -Exhaustively tested over the alphabet `{a, ., _}` up to length 6 -(`/tmp/probe_mangle.py`): **686 collisions in 1092 inputs.** Since a generated -identifier may only contain `[A-Za-z0-9_]`, `_` is the only available separator -and a *prefix escape* is required: - -```python -def codegen_name(tag: Tag) -> str: - return tag.replace("_", "_u").replace(".", "_d") - -def from_codegen_name(name: str) -> Tag: - return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) -``` - -Every `_` in the output is the first character of a two-character escape, so -decoding is unambiguous. Verified exhaustively over `{a, ., _, u, d}` up to -length 6 — **19530 inputs, 0 collisions, 0 round-trip failures**, every output a -valid identifier, including the adversarial `"_u"`, `"_d"` and `"a_ud.b"` -(`/tmp/probe_mangle2.py`). Cost: names grow (`mod.V2E.Local` → -`mod_dV2E_dLocal`), which is what gtfn's existing `TagDefinition.alias` mechanism -is for. - -**An inverse is needed too**, wherever generated names are parsed *back* into -dimensions: `dace/sdfg_args.py:25, 60-72` matches `gt_conn_(\S+)` and feeds the -result to `has_offset`. `codegen_name` must therefore be injective *and* have a -`from_codegen_name` partner (escape `__` → `____` before `.` → `__`). - -**A site that cannot be deferred: `gtir_to_sdfg_lambda.py:1155`.** It builds -`gtx_common.Dimension(offset, DimensionKind.LOCAL)` — a local dimension -synthesized from the **offset** tag, which in PR 2 is still a bare provider key -(`"V2E"`) that `resolve()` cannot import. Every DaCe unstructured shift passes -through it, so PR 2 is red on DaCe unless it is fixed there. The fix is local and -available: `conn_type` is already in scope (`:1134-1152`) and `:1135` already -asserts `conn_type.domain[1].kind == LOCAL`, so the line becomes -`offset_type = conn_type.domain[1]` (equivalently `conn_type.neighbor_dim`). -It is *necessary but not sufficient* for PR 1's `shift × tag≠localdim` DaCe cell: -that cell fails earlier, at `gtir_to_sdfg.py:842` -(`neighbor_table_types[dim.value]`, i.e. A4 on the connectivity *argument's* local -dim), before `:1155` is reached — and after the `:1155` fix, `:1371`/`:1455` would -reference `gt_conn_` while `:1104`/`:722` declare `gt_conn_`. So -**the DaCe shift cell stays in the skip matrix until PR 4**, where the -single-string choice makes both agree. (An earlier draft said PR 2; that would -leave PR 2 red on that cell.) - -**(c) The *offset* key is a second dotted name space, and it is mangled in PR 4, -not PR 2.** §1.3(b) covers *dimension* names only. When PR 4 makes the provider -key `cls.tag`, the **offset** string that flows through the IR -(`OffsetLiteral.value`, the provider key, `o` in the gtfn/DaCe connectivity -plumbing) becomes dotted too, and a different set of sites turns *it* into an -identifier. These are all **PR 4**: - -| Site | What breaks | -| --- | --- | -| `codegens/gtfn/itir_to_gtfn_ir.py:184` | `TagDefinition(name=Sym(id=offset_name))` → `SymbolName` regex | -| `codegens/gtfn/itir_to_gtfn_ir.py:490` | `SymRef(id=o)` for each connectivity → `SymbolRef` regex | -| `codegens/gtfn/codegen.py:147-148` | `visit_OffsetLiteral` emits `node.value` raw into C++ | -| `codegens/gtfn/gtfn_module.py:118, 132, 136` | `GENERATED_CONNECTIVITY_PARAM_PREFIX + name.lower()`, `generated::{name}_t` | -| `dace/sdfg_args.py:56` | `connectivity_identifier(name)` → `gt_conn_a.b.V2E`, an invalid SDFG array name | -| `dace/sdfg_args.py:60`, `dace/workflow/bindings.py:200, 286` | `is_connectivity_identifier` / `_parse_gt_connectivities` — the **inverse** direction, so `from_codegen_name` has *several* live consumers, not one | -| `dace/workflow/translation.py:61`, `dace/sdfg_callable.py:103`, `dace/program.py:156` | `connectivity_identifier(offset)` again, on the argument-binding path | -| `dace/lowering/gtir_to_sdfg_lambda.py:1104, 1371, 1455, 1727`, `gtir_to_sdfg.py:722` | the same identifier, consumed in the lowering | -| `runners/roundtrip.py:63` | `OffsetLiteral = as_fmt("{value}")` — emits the offset tag *raw as Python source*, into the program **body**; mangling `:176` alone still leaves `NameError: name 'tests' is not defined` | -| `dace/sdfg_args.py:83-84` | `_field_symbol`: `assert m[1] in offset_provider_type` — a *second* `from_codegen_name` consumer besides `:70` | -| `dace/lowering/gtir_to_sdfg_lambda.py:1892` | `visit_OffsetLiteral` → `SymbolExpr(node.value, INDEX_DTYPE)`, i.e. a dotted string used as a DaCe symbolic expression | -| `runners/roundtrip.py:152, 176` | collects offset-literal strings, then `f'{o} = offset("{o}")'` → `a.b.V2E = offset(...)` → `SyntaxError` | - -So `codegen_name` / `from_codegen_name` are introduced in PR 2 for dimensions and -**applied again in PR 4 for offsets**, at ~16 further sites. Two earlier claims -were wrong: that `from_codegen_name`'s only live consumer is in PR 2, and that the -DaCe surface is confined to `sdfg_args.py` and the lowering — the -argument-binding path (`workflow/translation.py`, `workflow/bindings.py`, -`sdfg_callable.py`, `program.py`) carries it too, in both directions. - -Because of (b) and (c), the review shortcut "diff PR 2 against #2844, the delta is -only identity" is **false**: #2844 needed none of this. Reviewers should expect a real -mangling layer on top of the identity delta. - -**`AxisLiteral.kind` becomes redundant** (the class carries it) — the `TODO` at -`iterator/ir.py:93`. Kept in PR 2, removed in PR 7, to keep PR 2's IR-expectation -churn to the `value` strings only. - -### 1.4 `LocalDimensionIndex` subclasses `DimensionIndex`; `DimensionBaseIndex` is dropped - -The proposal lists `DimensionBaseIndex` as a separate root with `DimensionIndex` -and `LocalDimensionIndex` as siblings. That does not survive contact with the -tree: `Dimension` is `type[DimensionIndex]`, eve validates a `type[X]` -annotation by `issubclass` (verified: a subclass passes, the base and an -unrelated class are both rejected), so sibling local dimensions would force -widening to `type[DimensionBaseIndex]` at `ts.DimensionType.dim`, -`ts.FieldType.dims`, `ConnectivityType.domain`, `Domain.__init__` and the `DimT` -/ `DimT_co` bounds — and would then accept local dimensions everywhere a primary -one is meant, which is the same looseness with extra ceremony. - -**Resolution**: `class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL)`. -`DimensionBaseIndex` is not introduced at all — one concept fewer, which is the -proposal's own stated goal. Where primary-only is required the check is -`dim.kind is not DimensionKind.LOCAL`, exactly as today. This also removes -#2844's deferral note ("a `DimensionBase` root, deferred until the requirements -of non-user-declarable dimensions are known") as a thing that needs resolving. - -**Verified**: all 38 sites in `src/` that discriminate a local dimension do so by -a **runtime `kind` check**, not by a static type distinction -(`transform_utils.py:65`, `type_deduction.py:460, 774`, -`custom_layout_allocators.py:171`, `past_to_itir.py:409`, `common.py:1168, 1336`, -`nd_array_field.py:972, 976`, `gtfn_module.py:91`, `embedded.py:922`, -`gtir_to_sdfg_types.py:76`, …). The tree already treats local dimensions as -`Dimension`s everywhere — `ConnectivityType.domain: tuple[Dimension, ...]` -includes the local one — so subclassing loses nothing it currently relies on, and -`Dims` (`tuple[Unpack[ShapeTs]]`, `common.py:57`) puts no bound on its members -either. - -**What subclassing does cost**, and the mitigation: every `DimensionIndex` -*bound* now statically admits a local dimension — -`NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]`, -`Staggered[D: DimensionIndex]` and `MultiDimensionIndex[D: DimensionIndex, *Ls]` -would all accept `V2E.Local` as their primary parameter. Each therefore gets a -runtime `kind is not DimensionKind.LOCAL` check in `__init_subclass__` / -`__class_getitem__`, and `LocalDimensionIndex.__init_subclass__` rejects an -explicit `kind=` other than `LOCAL`. This is the same runtime-check discipline -the tree already uses; the static gap is the price of the concept removed. - -**Deviation from the proposal; needs feeding back to the note.** - -### 1.5 `Staggered[D]` is required in PR 2, not PR 7 - -`flip_staggered` builds `Dimension(f"_Staggered{name}")` from a string -(`common.py:1452-1457`) and `is_staggered` tests `dim.value.startswith(prefix)` -(`:1447-1449`). #2844 routes both through the interning factory. With the -registry gone there is **no importable `_Staggered` type**, and a -dynamically created class would get the tag -`gt4py.next.common._Staggered`, so `is_staggered` is false and -`as_non_staggered` cannot recover the base dimension's module. Live dependents: -`test_staggered.py` (233 lines), `cases_utils.py:161` -(`KHalfDim = flip_staggered(KDim)`), gtfn `_add_staggered_aliases` -(`itir_to_gtfn_ir.py:203-215`), DaCe `get_map_variable` -(`gtir_to_sdfg_utils.py:52`), `type_synthesizer`, `test_common.py`, -`test_domain_utils.py`. - -So PR 2 is **not green** without `Staggered[D]`. It is Cartesian-only and does -not depend on the connectivity layer, so it moves into PR 2. - -**The obvious mechanism does not work.** A PEP 695 generic -`class Staggered[D: DimensionIndex](DimensionIndex)` makes `Staggered[KDim]` a -`typing._GenericAlias`, **not a class** (verified, `/tmp/probe_staggered.py`): - -``` -type(Staggered[KDim]) -> -isinstance(Staggered[KDim], type)-> False -issubclass(Staggered[KDim], ...) -> TypeError: issubclass() arg 1 must be a class -Staggered[KDim].tag -> '__main__.Staggered' # KDim is gone -``` - -So it fails eve's `type[DimensionIndex]` validation and its tag cannot name the -base dimension — it is not a `Dimension` at all. - -**The mechanism that does work** (verified, `/tmp/probe_staggered3.py`: runs -correctly and is **0 errors under both `mypy --strict` and pyright 1.1.414** on -3.12) is a metaclass `__getitem__` that *builds and interns a real class*, paired -with a `TYPE_CHECKING` declaration so checkers still see an ordinary generic: - -```python -class StaggeredMeta(DimensionMeta): - def __getitem__(cls, base: Dimension) -> Dimension: - if base not in _staggered_cache: - _staggered_cache[base] = StaggeredMeta( - f"Staggered[{base.__name__}]", - (cls,), # NOT (cls, base) -- see below - {"_tag": f"{cls.__module__}.{cls.__qualname__}[{base.tag}]", - "kind": base.kind, "base": base, "__slots__": ()}, - ) - return _staggered_cache[base] - -if TYPE_CHECKING: - class Staggered[D: DimensionIndex](DimensionIndex): - base: ClassVar[Dimension] -else: - class Staggered(DimensionIndex, metaclass=StaggeredMeta): - __slots__ = () - base: ClassVar[Dimension] -``` - -Verified properties of `Staggered[KDim]`: it *is* a class; -`tag == "gt4py.next.common.Staggered[]"`; `kind` is -inherited from the base; `issubclass(_, DimensionIndex)` and -`issubclass(_, Staggered)` hold; it is instantiable as an index; and -`Staggered[KDim] is Staggered[KDim]`, so identity is stable. `Staggered[KDim]` in -an annotation and inside `Field[Dims[Staggered[KDim]], float]` are both accepted -by both checkers. - -- **Bases are `(cls,)`, not `(cls, base)`.** Inheriting from the base dimension - would make `issubclass(Staggered[KDim], KDim)` true, i.e. `KHalfDim` would be - accepted everywhere `KDim` is required. It is a *different* dimension; only - `kind` is inherited, copied explicitly into the namespace. -- `is_staggered(dim)` becomes **`"base" in dim.__dict__`**, not - `issubclass(dim, Staggered)`, and `as_non_staggered(dim)` becomes `dim.base`. - Two runtime facts force this: `issubclass(Staggered, Staggered)` is true for - the bare base, which has no `base`; and `Staggered[KDim]` is *subclassable* - (`class KHalf2(Staggered[KDim])` yields a second, un-interned staggered-K type - with tag `.KHalf2`). `Staggered.__init_subclass__` therefore rejects - any subclass the metaclass did not create, so the interned form is the only - one. Still structural — no string sniffing. -- **The guards were verified, including the escape routes** - (`/tmp/probe_staggered_guards.py`). All four are blocked: - `class KHalf2(Staggered[KDim])`, `class X(Staggered)`, a direct - `StaggeredMeta("Y", (Staggered,), {})`, and the double subscript - `Staggered[KDim][KDim]`. The `copyreg` fallback round-trips the bare - `Staggered` by reference, the parametrized class with identity preserved, and - instances. Implementation note: gate `__init_subclass__` on a **namespace - marker** the metaclass sets (`"_tag" in cls.__dict__`), not on a module-level - "currently building" flag — the flag works but is not thread-safe, and - compilation runs in worker processes and threads. Three further honest limits: - the guard defends against **accidental** subclassing only — a deliberate - `StaggeredMeta("Forged", (Staggered,), {...marker})` or `types.new_class` can - still forge a same-`tag`, non-identical type (as it can for any class); - `Staggered[Staggered[KDim]]` must be rejected explicitly by testing - `"base" in base.__dict__` in `__getitem__`, or it nests and pickles happily; and - a hand-built `copyreg` payload such as `(_make_staggered, (int,))` should raise - a `TypeError` naming the offending base rather than an `AttributeError`. -- `resolve` gains the `[]` grammar: it parses the brackets and - evaluates `Staggered[resolve(inner)]`, which hits the same intern cache, so a - staggered dimension round-trips through the IR to the *same* class object. -- **A narrow `copyreg` is required after all.** `Staggered[KDim]`'s - `__qualname__` is `Staggered[KDim]`, which `pickle.save_global` cannot look up: - `PicklingError: Can't pickle : attribute lookup - Staggered[KDim] on … failed` (verified, `/tmp/probe_staggered_pickle.py`). A - `copyreg.pickle(StaggeredMeta, lambda cls: (_make_staggered, (cls.base,)))` - fixes it *and preserves identity*, because the reconstructor goes back through - the intern cache. **But it must guard the bare base**: `type(Staggered) is - StaggeredMeta` too, so a reducer that unconditionally reads `cls.base` fails on - `Staggered` itself with `AttributeError: type object 'Staggered' has no - attribute 'base'` (verified — an earlier draft of this section claimed the - registration "captures only parametrized dimensions", which is false). The - reducer therefore falls back to by-reference pickling when - `"base" not in cls.__dict__`. It never captures a plain dimension - (`type(KDim) is DimensionMeta`). This is materially narrower than - #2844's blanket registration on `DimensionMeta` — a parametrized type needs a - reconstructor for the same reason `typing` aliases do — but §0's "`copyreg` - dropped" row is only true of the blanket form, and the ADR must say so. -- **Two honest costs.** (i) `_staggered_cache` is a cache, and the proposal's - headline is that the *name-keyed* registry goes away. The difference is real but must be - stated: it is keyed by a *dimension class*, is internal, and is memoization of - a type constructor (as `typing`'s own subscription cache is), not interning of - user-authored name strings — nothing resolves a user string through it. - (ii) the `TYPE_CHECKING` split means the static and runtime definitions can - drift; a unit test must assert the runtime facts the static form does not - express (real class, `issubclass` against `Staggered` but *not* against the - base, tag shape, interning). -- Supersedes ADR 0026, recorded in the PR-2 ADR. - -## 2. The PR stack - -Branches follow the repo's stacked convention, `connectivities-as-types--`, -each based on its predecessor, all targeting `main`. PR titles are Conventional -Commits (squash-merge lands the title). - ---- - -### PR 1 — `fix[next]: lower unstructured shifts with the offset's own tag` - -**Independent of the rest of the stack; lands first, on its own merit.** - -`foast_to_gtir._visit_shift` emits the **Python variable name** as the IR shift -tag (`foast_to_gtir.py:305` `offset_name.id`, `:331` `str(offset_name)`), because -`ts.OffsetType` does not carry the tag. Embedded execution keys on -`FieldOffset.value`. So the same program needs a *different* provider key -depending on the backend — confirmed by running it on v1.2.2: - -``` -MyOff = FieldOffset("TAGNAME", ...) -embedded: {"TAGNAME": conn} OK ; {"MyOff": conn} -> KeyError 'TAGNAME' -roundtrip: {"MyOff": conn} OK ; {"TAGNAME": conn} -> KeyError 'MyOff' -``` - -**Change** - -- `ts.OffsetType` gains **`tag: Optional[Tag] = None`** — *not* a required field. - `type_deduction.py:709` builds `ts.OffsetType(source=conn.codomain, - target=(conn.domain_dim,))` from `IDim + 1`, a `CartesianConnectivity` that has - no tag at all; making `tag` required breaks it. -- `FieldOffset.__gt_type__` fills it (`fbuiltins.py:485`). -- `type_deduction.py:464`, which rebuilds an `OffsetType` when `Off[1]` drops the - local dimension, must **propagate** the tag. -- `foast_to_gtir._visit_shift`: the `Subscript` branch and the bare `Name` branch - use `arg.type.tag`, asserting non-`None` (both are unstructured paths, where a - tag always exists). - -**Tests.** `tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py` -today covers exactly `a(Off[1])` on `GTFN_CPU`. Extend to -{shift, `neighbor_sum`} × {embedded, roundtrip, gtfn, dace} × {tag≠varname, -tag≠local-dim-name}. - -**The matrix is not uniform, and a blanket `xfail` will not do.** `xfail_strict = true` -(`pyproject.toml:323`), and measured behaviour on v1.2.2 is: - -| case | embedded | roundtrip | gtfn | dace | -| --- | --- | --- | --- | --- | -| shift, tag≠varname | pass | pass | pass | pass | -| shift, tag≠localdim | pass | pass | pass | **fail** `KeyError` (`gtir_to_sdfg_lambda.py:1155` synthesizes the local dim from the tag) | -| `neighbor_sum`, tag≠localdim | **fail** | **pass** | **fail** | **fail** | - -So the first draft's acceptance criterion ("shift cells pass on all four -backends") is unreachable before the backend work, and a strict blanket `xfail` -would XPASS on roundtrip. **Fix**: add a per-backend skip matrix entry in -`tests/next_tests/definitions.py` (a new `USES_*` marker) covering exactly the -failing cells, roundtrip excluded. **They are removed in two steps**: the -`shift × tag≠localdim` DaCe cell and the three `neighbor_sum × tag≠localdim` cells -all in **PR 4**, where the single-string choice makes A3/A4 vacuous — *not* in -PR 5, and *not* the DaCe cell in PR 2 (the `:1155` fix there is necessary but not -sufficient; see §1.3(a)). The gtfn `neighbor_sum` -failure is now confirmed **by running it**; the proposal had it only "by -reading". - -**No CHANGELOG entry.** Verified against the history: `CHANGELOG.md` is touched -*only* by release PRs (`git log -- CHANGELOG.md` is release commits exclusively, -and nothing between `b3c53fa7e` and `upstream/main` touches it). The behaviour -change — which key a compiled backend requires when tag ≠ variable name — belongs -in the PR description, and reaches the changelog when the release PR is cut. Two -earlier drafts of this plan said otherwise, including for PR 6's breaking change. - -**ICON4Py is unaffected by PR 1**: all 16 `FieldOffset` variable names equal -their tags (`model/common/src/icon4py/model/common/dimension.py:33-48`). - -**Acceptance**: `nox -s test_next` green; every cell in the matrix either passes -or is covered by the documented skip matrix. - ---- - -### PR 2 — `feat[next]: a concrete Dimension is a class, identified by its qualified name` - -The #2844 core with the identity divergences of §0, **plus** the mangling layer -of §1.3(b) and `Staggered[D]` of §1.5 — both of which #2844 did not need and -without which this PR cannot be green. Large and largely mechanical. - -**`src/gt4py/next/common.py`** - -```python -class DimensionMeta(type): - kind: DimensionKind - __hash__ = type.__hash__ # §1.0 — mandatory, not optional - @property - def tag(cls) -> Tag: ... # f"{cls.__module__}.{cls.__qualname__}" - # operators as in #2844; __eq__/__ne__ keep only the IntegralScalar overload - # (I == 5 -> Domain); the dim-vs-dim branch is `is`. - -class DimensionIndex(metaclass=DimensionMeta): - __slots__ = ("value",) - kind: ClassVar[DimensionKind] = DimensionKind.HORIZONTAL - def __init_subclass__(cls, /, kind=None, **kw): ... - -# Staggered: an interning metaclass + TYPE_CHECKING split, NOT a PEP 695 -# generic -- see §1.5, where the generic form is shown to be unworkable. - -def resolve(tag: Tag) -> Dimension: ... # memoized; [] grammar -def codegen_name(tag: Tag) -> str: ... # "_" -> "_u", "." -> "_d" (§1.3(d)) -def from_codegen_name(name: str) -> Tag: ... # the inverse, `_([ud])` -> `_` / `.` - -type Dimension = type[DimensionIndex] -``` - -- `tag` is a metaclass **property**, so it cannot drift from the type. This makes - a class-body `tag = "C2E"` a **silent no-op** (verified: `C2EDim.tag` stays - `"__main__.C2EDim"` even with `tag = "C2E"` in the body) — and that pattern is - exactly what ICON4Py and #2845 use to rename. `__init_subclass__` therefore - **raises** on `"tag" in cls.__dict__`, naming the class and pointing at the - rename path. -- `__init_subclass__` also rejects `"" in cls.__qualname__`. Neither - necessary nor sufficient (`type("Dyn", ...)` in a function passes; a `del`'d - class passes) — the authoritative check stays pickle's own `save_global`. -- `resolve` raises a `ValueError` naming the tag and the failing import, per - CODING_GUIDELINES. - -**Removals**: `NamedIndex`; `_DimA`..`_AnyDim` and the dimension half of -`mypy_plugin.py`; `_DIMENSION_REGISTRY`; `copyreg`; the `fingerprinting.py` -deconstructor; `_STAGGERED_PREFIX` and its string sniffing. - -**Migration**. Every `Dimension("X")` becomes `class X(DimensionIndex): ...` at -module level. Verified counts: 333 `Dimension("` declarations in `tests/`, of -which **133 are function-local across 15 files** and must move to module level; -a dimension *named* `"I"` is declared 46 times across **17** files (the first -draft said 45 files — that was the proposal's *`IDim` file* count, a different -number). Docs, workshop notebooks and `examples/` are included because -`test_examples` executes them; notebook *code* cells only, stored outputs -untouched (they hold recorded tracebacks that must keep naming the symbols that -produced them). - -**IR expectation churn**: 36 `AxisLiteral` and 37 `OffsetLiteral` occurrences in -`tests/`, most already computed from `dim.value`. `test_pretty_roundtrip.py` and -the gtfn/DaCe snapshot tests hold the hardcoded names. - -**Do the sweep with a codemod script, not agent fan-out.** A previous attempt at -agent fan-out on a large mechanical rewrite in this repo died mid-file on the -rate limit and left the tree inconsistent; a script did all 57 files uniformly. - -**ADR 0028** (the directory ends at 0027): nominal identity; the module-level -declaration requirement; `Staggered[D]` superseding ADR 0026; that cache -fingerprints now shift when a declaration moves module (a consequence for ADR -0023, not a reversal); that `resolve()` imports modules named in the IR, which is -the same trust level as `pickle` loading a class by reference. - -**Documented limitation**: interactive `__main__` (REPL, notebooks, `python -c`) -cannot be resolved. `spawn` compile workers re-execute the main *script* as -`__mp_main__`, so file-based `__main__` resolves provided the script has the -`if __name__ == "__main__":` guard the pool already requires. - -**ICON4Py migration script is a PR-2 deliverable, not PR 6.** All 15 local -dimensions and `KDim`/`EdgeDim`/`CellDim`/`VertexDim` have variable name ≠ tag -(`EdgeDim = Dimension("Edge")`), so PR 2 changes every generated symbol and every -cache key downstream. - -**Acceptance**: `nox -s test_next` on **3.12, 3.13 and 3.14** (the `typing` -subscription cache behaves differently per interpreter and this change moves -exactly that behaviour), then `test_eve`, `test_storage`, `test_cartesian`, -`test_examples`; `uv run mypy src/`; `uv run pyright`; `uv run tach check`; -`uv run pre-commit run -a`. One at a time, pytest capped at `-n 4`. - ---- - -### PR 3 — `feat[next]: NeighborConnectivity declarations and local dimensions that know their owner` - -**Purely additive**: new concepts next to `FieldOffset`, nothing removed, no -behaviour change, no test churn. - -```python -class LocalDimensionIndex(DimensionIndex, kind=DimensionKind.LOCAL): # §1.4 - owner: ClassVar[type[NeighborConnectivity] | None] = None - max_neighbors: ClassVar[int | None] = None - min_neighbors: ClassVar[int | None] = None - def __init_subclass__(cls, *, size: int | None = None, **kw): ... - -class ConnectivityMeta(type): # §1.0 for __hash__ and __getitem__ - @property - def tag(cls) -> Tag: ... - def __call__(cls, *a, **kw) -> NoReturn: ... - -class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( - metaclass=ConnectivityMeta -): - Local: ClassVar[type[LocalDimensionIndex]] - def __init_subclass__(cls, *, max_neighbors=None, min_neighbors=None, **kw): ... -``` - -- `__init_subclass__` asserts `"Local" in cls.__dict__` and that it subclasses - `LocalDimensionIndex`, then sets `Local.owner = cls` and copies the counts. A - missing `Local` is a `TypeError` at class creation naming the class. -- **Owner-less** locals: `class LsqUnk(LocalDimensionIndex, size=3)` — `owner is - None`, `min == max == size`, never in the provider. ICON4Py's `LsqUnkDim` and - `RBFDimension` need this: they index no table but need sparse storage and - layout. -- Counts are **optional class keywords**, not type parameters (Python has no - integer type parameters and nothing static needs the count). Declared ⇒ a - constraint the table must satisfy. Undeclared ⇒ completed at bind time, from - the table in the JIT flow or from the `NeighborConnectivityType` already passed - through `connectivities=` (`ffront/decorator.py:188-208`) in the AOT flow. - Not static-only because `fvm_nabla_setup.py:99` sizes `V2E` from the atlas - mesh, and ICON skip-value presence is configuration-dependent (`icon.py:130` — - pentagons have skip values on the icosahedron, not the torus). -- **Bind-time validation**, one function replacing constraints A6–A8: shape - `(n, max_neighbors)`, integral dtype, skip values present iff - `min_neighbors < max_neighbors`, `domain[0] is Origin`, `codomain is Codomain`. - -**No provider bridge.** The first draft proposed a dual-keyed -(`Tag | type[NeighborConnectivity]`) provider here. Dropped: the provider is -accessed **directly, not through `get_offset`, at 19 sites in 12 `src/` files** -(`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, -`compiled_program.py`, `pass_manager.py`, …) despite the note at -`common.py:1174`, so a bridge would be both invasive and — since nothing would -exercise class keys — untested. Class-keyed providers land in one place, PR 4. - -**Typing tests**: `typing_probe.py` / `probe_local.py` / `probe_meta_getitem3.py` -become real coverage — `typing_tests/test_next.yaml` cases for -`Field[Dims[V, V2E.Local], float]`, a `TypeVar` bound to `LocalDimensionIndex`, -the negative cross-connectivity case, and the `NC[V, E]`-in-bases case of §1.0; -plus the pyright variants. - -**Acceptance**: full suite green with no behaviour change; new unit tests for -declaration errors, owner wiring, owner-less locals and bind-time validation. - ---- - -### PR 4 — `feat[next]: declare connectivities as classes; FieldOffset derived from them` - -**The ordering fix.** Revision 2 put class-keyed providers here and the -declaration migration in PR 6. That cannot be green: `Local.owner` only exists if -the user declared a `NeighborConnectivity` class, and a `FieldOffset` written the -old way (`FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim))`) has no class -to point at — so neither the backend work nor a class-keyed provider has anything -to resolve. The declaration migration must come **first**, and the provider key -must stay a string until the backends are through. - -- The **unstructured** `FieldOffset` is **derived**, not authored: - `FieldOffset.from_connectivity(V2E)` (or `V2E.__gt_offset__()`), which fills - `source = Codomain`, `target = (Origin, V2E.Local)` and — critically — - **`value = V2E.Local.tag`, the *local dimension's* tag, not `V2E.tag`.** -- **Why the local dimension's tag and not the connectivity's.** An earlier draft - used `V2E.tag` and claimed PR 4 was green. It is not: A3 and A4 key the provider - on the **local dimension's** name, at `nd_array_field.py:983` - (`get_offset(provider, axis.value)`, whose in-tree comment is literally - `# assumes offset and local dimension have same name`), `unroll_reduce.py:47` - (`arg.type.offset_type.value`), `gtfn_module.py:95`, `gtir_to_sdfg.py:581, 842`, - `iterator/embedded.py:954, 1519`, and - `gtir_to_sdfg_lambda.py:1371, 1455` (`connectivity_identifier(offset_type.value)`). - Today the `V2EDim = Dimension("V2E")` convention makes that string equal to the - offset tag; PR 4 deletes the convention tree-wide, while the `owner` lookup that - replaces it is PR 5. With `value = V2E.tag` every reduction and every sparse-field - argument would break on embedded, gtfn and DaCe simultaneously — the round-1 - matrix row (`neighbor_sum`, tag≠localdim: embedded/gtfn/DaCe fail) would become - the tree's universal state. - Choosing `V2E.Local.tag` instead makes **all four** of A1, A3, A4 and A5 vacuous - at once, because there is then exactly *one* string and the class produces it. - It also makes the #1789 branch at `itir_to_gtfn_ir.py:181-190` - (`if offset_name != connectivity_type.neighbor_dim.value`) dead already in PR 4. - This is preferable to the alternatives — fusing PR 5 into PR 4, or a transient - `owner`-based fallback inside `get_offset` — because it needs no scaffolding: - the string is simply picked correctly, and PR 5 then removes the dependence on a - string at all. -- **The Cartesian `FieldOffset` constructor stays in PR 4.** An earlier draft said - "every declaration becomes a class", which is wrong: `Ioff`, `Koff` and - `EdgeOffset` (`cases_utils.py:163-169`, e.g. - `Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,))`) and ICON4Py's - `Koff`/`KHalfOff` (`dimension.py:47-48`) are single-target and have no - `NeighborConnectivity` to derive from, and their only remaining consumer — - `as_offset` — does not change until PR 6. Restricting the public constructor to - the single-target form keeps PR 4 green; it disappears with `as_offset` in - PR 6. -- **Providers stay keyed on `Tag`**, now `V2E.Local.tag`. Nothing about the key - *mechanism* changes yet, so the 19 direct-access sites are untouched. A1, A3, A4 - and A5 are all dead at this point — there is one string, and the class produces - it. -- `ts.OffsetType` → `ConnectivityType`, produced by the class - (`type_specifications.py:74` TODO). `type_info.py:637, 858` gate `a(V2E)` - deduction on `ts.OffsetType` and follow. `Field.premap` / `Field.__call__` - unions widen to `Connectivity | type[NeighborConnectivity]` (`common.py:785, - 791-794`, `1313-1316`). -- **Test-tree migration lands here**: `toy_connectivity.py`, `cases_utils.py`, - `fvm_nabla_setup.py` are the fixture modules everything imports; 37 - `FieldOffset` sites, 42 `DimensionKind.LOCAL` sites. String provider keys keep - working because they are `cls.tag` — but the tags are now *qualified*, so the - 111 ITIR-level string-offset occurrences in 15 files (`im.shift("V2E")`, - `neighbors("…")`, `OffsetLiteral(value="…")`, string-keyed providers) are - rewritten to `V2E.Local.tag` here rather than in PR 6. -- **ICON4Py**: this is the release-visible declaration change. The migration - script written in PR 2 is extended. - ---- - -### PR 5 — `refactor[next]: backends resolve connectivities through the local dimension's owner` - -Where A3, A4 and A5 dissolve and PR 1's skip-matrix entries are removed. Green -while providers are still string-keyed, because a backend goes -`local_dim.owner` → `owner.tag` → the existing lookup: the *identity* question is -answered by the owner pointer, and the key is still a string. - -**The owner-less case must be handled, not assumed away.** At PR 5 `_CONST_DIM` -is still a plain `DimensionIndex(kind=LOCAL)` (it becomes `ConstList` only in -PR 7), and `LsqUnk`-style local axes have `owner is None` by design. Every -converted site reads `getattr(dim, "owner", None)` and falls through when it is -`None`. Two sites already guard by accident — DaCe compares against `_CONST_DIM` -first (`gtir_to_sdfg_lambda.py:1314`) and `unroll_reduce` filters -`offset_type is None` — but `gtfn_module.py:91-98` and `nd_array_field.py:981` -have **no** guard, and ICON4Py never exercises the case -(`test_icon.py:220`), so the gap would not show up downstream. - -`unroll_reduce.py:47` (reads `arg.type.offset_type`, which is the local -`Dimension` — now a `LocalDimensionIndex` carrying `owner`, which is exactly the -back-pointer it lacked), `gtfn_module.py:95, 118, 132`, `itir_to_gtfn_ir.py` -(including the `#1789` `offset_name != neighbor_dim.value` branch at `:181-190`, -which becomes dead and goes), `gtir_to_sdfg.py:581`, -`gtir_to_sdfg_lambda.py:766-770, 1155` (the `Dimension(offset, LOCAL)` synthesis -goes), `nd_array_field.py:981-985` (and its -`# assumes offset and local dimension have same name` comment), -`iterator/embedded.py`, and `runners/roundtrip.py`. - ---- - -### PR 6 — `feat[next]!: class-keyed offset providers; remove FieldOffset and the string offset API` - -The only breaking PR, and now the only one that touches the provider key. - -- `OffsetProvider*` become - `Mapping[type[NeighborConnectivity], NeighborTable]`; `get_offset` keys on the - class. **Note the `.owner` hop**: because PR 4 made the IR offset tag the *local - dimension's* tag, `resolve(OffsetLiteral.value)` yields a `LocalDimensionIndex`, - not the connectivity — so the class-key lookup is - `resolve(tag).owner`. (The alternative is to switch the IR tag to `V2E.tag` in - this PR; the `.owner` hop is cheaper and keeps the IR stable.) The **19 direct-access sites in 12 `src/` files** - (`itir_to_gtfn_ir.py:181`, `gtfn_module.py:106`, `sdfg_args.py:63-72`, - `compiled_program.py`, `pass_manager.py`, …) are converted here, and - `common.py:1174`'s "all accesses should go through `get_offset`" either becomes - true or the note goes. `hash_offset_provider_items_by_id` and the - `fingerprinting` dict handling already tolerate class keys once §1.0's - `__hash__` is in place. -- **Removals**: `FieldOffset` entirely (both forms); `runtime.Offset` as its base - (`fbuiltins.py:467-470` TODO); `iterator/runtime.offset("...")` (12 sites in 6 - files, plus `tracing.py:161-162`); the `V2EDim`-next-to-`V2E` convention; - `embedded/context.py` string plumbing; the `gt4py.next.__init__` exports at - `:47, 140`. -- **`as_offset` changes in the same PR.** It is why the Cartesian `FieldOffset` - form cannot go alone: `ffront/experimental.py:17` + - `type_deduction.py:956-967` require one. New signature - `as_offset(KDim, field)`. Used in 5 test modules, the `Ioff`/`Koff`/`EdgeOffset` - fixtures at `cases_utils.py:163-169`, and **40 non-test call sites in - ICON4Py**. -- `transform_utils.py:50-77` and `past_to_itir.py:77` deduce grid type from the - provider and follow. -- **Accepted double churn**: the ~26 `offset_provider={...}` literals are rewritten - twice — `{V2E.Local.tag: t}` in PR 4, `{V2E: t}` here. The alternative is fusing PR 4 - and PR 6, which loses the green boundary. The ITIR string sites do *not* churn - twice: `im.shift(V2E.Local.tag)` written in PR 4 stays correct. - -**ADR 0029**: the connectivities-as-types record — `FieldOffset` removed, the -class-keyed provider, superseding the `FieldOffset` part of ADR 0019. - -**Breaking-change communication**: the PR title carries the Conventional Commits -`!` marker and the ADR records the removal; the changelog entry is written by the -release PR, not here (see PR 1). No deprecation window — an explicit decision: -ICON4Py's provider keys are bare names, so no import-based shim could have -resolved them. - -**Acceptance**: full suite; `test_fvm_nabla` and `test_icon_like_scan` are the -integration canaries. A before/after run of one gtfn and one DaCe program -checking generated-code equivalence modulo names. - -### PR 7 — `refactor[next]: ConstList replaces _CONST_DIM; AxisLiteral drops kind` - -- `_CONST_DIM` (`iterator/embedded.py:220` and - `dace/lowering/gtir_to_sdfg_lambda.py` — two separate declarations, each - internally consistent; 14 references in total) becomes the owner-less - `ConstList(LocalDimensionIndex, size=1)`, generalizing the magic name from size - 1 to size *n*. -- `AxisLiteral.kind` removed (`iterator/ir.py:93` TODO), now that every tag - resolves to a class carrying its kind. - ---- - -### PR 8 — `refactor[next]: MultiDimensionIndex and typed embedded positions` - -- `MultiDimensionIndex[D: DimensionIndex, *Ls]` as the index type of a sparse - position and the domain index of a `NeighborTable`. `*Ls` is unconstrained - because `TypeVarTuple` cannot carry a bound; `__init_subclass__` checks at - runtime what the checker cannot. -- `iterator/embedded.py` positions keyed by dimension types instead of name - strings (`embedded.py:574-576`, `597-616`, `941-950`); `SparseTag` removed. - Constraint A9 dissolves. -- Nothing else depends on this; it is last for that reason. - ---- - -## 3. Constraint ledger - -| # | Constraint | Retired by | -| --- | --- | --- | -| A1 | `FieldOffset.value` == provider key | PR 4 (one declaration produces both) | -| — | *all four string-equality constraints below become vacuous in PR 4*, because the class emits a single string (`V2E.Local.tag`); PR 5 removes the dependence on a string at all | PR 4 / PR 5 | -| A2 | Python variable name == provider key | **PR 1** | -| A3 | local dim name == provider key (reductions) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | -| A4 | local dim name == provider key (sparse args) | PR 4 vacuous; mechanism removed in PR 5 (`Local.owner`) | -| A5 | `FieldOffset.value` == local dim name | PR 4 (one declaration) | -| A6 | `target[-1]` == connectivity `neighbor_dim` | PR 3 (bind-time check) | -| A7 | `FieldOffset.source` == `codomain` | PR 3 (bind-time check) | -| A8 | `target[0]` == `domain[0]` | PR 3 (bind-time check) | -| A9 | dim name is the iterator-position dict key | PR 8 | -| A10 | dim name round-trips through `AxisLiteral` | structural; PR 2 qualifies it, PR 7 drops `kind` | -| F5/F8/F9 | codegen name formats | PR 2 (`codegen_name` + inverse) | -| S6 | `as_offset` needs a Cartesian `FieldOffset` | PR 6 | -| — | string-keyed provider | PR 6 | - -## 4. Risks - -1. **PR 2 size, and it is not purely mechanical.** ~150 files of migration *plus* - a name-mangling layer and `Staggered[D]`. It cannot be reviewed as "the - #2844 diff plus identity". Mitigation: land the mangling layer and - `Staggered[D]` as reviewable commits *within* the PR, ordered before the - sweep, so the mechanical part is a separate commit. -2. **Function-local dummy dimensions.** 133 declarations in 15 test files must - move to module level. Under nominal identity, two same-named locals that were - silently the same dimension become distinct — each resulting failure is a - real finding, not churn. -3. **`typing` subscription caching.** Under `(name, kind)` equality, - `Field[Dims[I]] is Field[Dims[I2]]` aliases for two distinct same-named - classes — a known residual of the #2844 design, and an argument *for* this - stack. The cache behaves differently per interpreter: verify on 3.12, 3.13 - and 3.14 **via nox**, not `uv run pytest`. -4. **Fingerprint/cache invalidation.** Moving a declaration between modules now - invalidates compiled artifacts. Intended; CHANGELOG + ADR line. -5. **Naming not yet converged** with `havogt/dependent-local-dimensions` - (`Origin`/`Codomain` vs `source_dim`/`neighbor_dim`; `Local` vs `Dim`; - `min_neighbors` vs `has_skip_values`). PR 3 fixes public names. Converge - before PR 3 is *opened*. -6. **`V2E.Local` vs `Local[V2E]`.** The chain proposals' encodings subscript - `Local`, and a `TypeVar` cannot be subscripted for a nested attribute. Their - semantics are unaffected; their static encoding needs rewriting to `C.Local` - plus a protocol for the generic hop-stack case. Knowledge-repo concern, not a - gt4py blocker. -7. **CSCS GPU CI is flaky and opaque.** All jobs failing at the same second means - infrastructure; `cscs-ci run default` as a PR comment reruns it. #2844's CI is - green except that job. -8. **Two mangling passes, two PRs.** `codegen_name` is applied to dimension names - in PR 2 and to offset names in PR 4, at ~19 sites total, several of which - (`Sym`/`SymRef` construction, DaCe array names) fail *loudly* and several of - which (C++ emission, `name.lower()`) fail only in the generated artifact. - Both PRs need a test that a qualified tag survives a real gtfn and a real - DaCe compile, not just lowering. -9. **`resolve()` on the inference hot path** (`inference.py:464`, once per - `AxisLiteral`). Must be memoized from the start, and the memo must be keyed - so a reloaded module does not return a stale class. - -## 4b. Work the earlier drafts did not mention - -- **Public exports.** `gt4py.next.__init__` must export `NeighborConnectivity`, - `LocalDimensionIndex`, `Staggered` and `resolve` (PR 2 for the dimension half, - PR 3 for the connectivity half), and drop `FieldOffset` / `offset` at `:47, 140` - in PR 6. -- **`type_translation.from_value(V2E)`** works only because the - `hasattr(value, "__gt_type__")` branch at `type_translation.py:328` is tested - *before* the `DimensionMeta` branch. That ordering is load-bearing under this - design and currently untested — PR 3 adds a unit test pinning it. -- **`pyright` is not yet a dependency.** `uv run pyright` appears throughout §5 - but pyright is absent from `pyproject.toml` on `main`; #2845 is what adds it. - PR 2 must explicitly fold in #2845's `typing_exports` / pyright dependency-group - change, or §5's pyright step is not runnable. -- **`test_examples` belongs to PR 4 too.** §5 lists it for PR 2 and PR 6; the - docs and notebooks use `FieldOffset`, so PR 4's declaration migration touches - them and must run it. - -## 5. Verification - -Per PR, in this order, **one at a time** on the shared machine, pytest capped at -`-n 4`: - -``` -uv run pre-commit run -a # ruff, mypy, tach, license headers -uv run pyright # static checks the mypy plugin no longer fakes -uv run nox -s "test_next-3.12(...)" # then 3.13, 3.14 for PR 2 -uv run nox -s test_eve test_storage test_cartesian test_examples # PR 2, PR 6 -``` - -Test-first where behaviour changes, per AGENTS.md: PR 1's regression matrix, PR -3's declaration-error and bind-validation units, and PR 4's provider-key tests -are written before the implementation they cover. - -## 6. Feedback owed to the knowledge-repo note - -- `NeighborConnectivity` is not a `Connectivity`; the sketch's base line is wrong - (§1.1) — this closes Open Q6. -- `DimensionBaseIndex` should be dropped; `LocalDimensionIndex` subclasses - `DimensionIndex` (§1.4). -- Open Q2 is closed: metaclass discovery, base `ClassVar` *annotation*, explicit - nested class — with the precise limit of what that buys statically (§1.2). -- `Staggered[D]` is not a late step; it is a precondition for removing the - name-keyed registry (§1.5) — and it cannot be a PEP 695 generic, needs an - identity-keyed intern cache, and needs a narrow `copyreg`. The note's claim - that `copyreg` disappears entirely is therefore too strong. -- The note's "five name spaces" analysis should record that under qualified tags - there are **two** dotted name spaces reaching codegen — dimension tags and - offset tags — each needing its own mangling pass (§1.3(b) and (c)). -- The note's §Staging step 2 ("`NeighborConnectivity` … object-keyed provider; - `FieldOffset` and string keys removed outright") bundles three changes that - must be separated to stay green: the declaration migration has to precede the - backend work (because `Local.owner` only exists once classes are declared), and - the provider *key* has to stay a string until the backends resolve through the - owner. See PR 4/5/6. -- The staging in the note's §Staging (steps 0–8) is superseded by §2 here; in - particular step 4 ("backends, one file at a time") cannot follow step 2, since - the provider key change and the backend lookups are separable but the mangling - layer is needed at the *dimension* step. - -## 7. Open, non-blocking - -- Naming convergence (risk 5). -- Whether a same-`__name__` collision warning is useful or noise (note Q5). -- Whether interactive `__main__` should be detected with a fallback to in-process - compilation and a warning, or merely documented (note Q3). Plan assumes - documented. -- `AxisLiteral.dim` instead of `AxisLiteral.value` — a follow-up after PR 8. From 03cb2f2de25789123f6156714d133705b8b219a4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 13:30:52 +0200 Subject: [PATCH 06/20] wip[next]: offset-name mangling and DaCe/gtfn fixes; all unit tests green `tests/next_tests/unit_tests` green on every backend: 1760 passed without DaCe, 451 with. Integration suite not yet re-run. The main fix is a consequence of the previous commit that it did not follow through. Pulling the one-string-per-connectivity invariant into this PR makes the *offset* keys qualified here too, so the offset-name mangling the plan scheduled for PR 4 is needed now: * gtfn: connectivity parameter names and `generated::_t` keys, the `SymRef` to each connectivity, and `axis_literal` literals. * DaCe: `connectivity_identifier` mangles, and all three places that parse an identifier back (`is_connectivity_identifier`, `_field_symbol`, and the generated binding code) unmangle -- the binding uses the mangled name as a Python variable but must look the table up by the real tag. * DaCe `visit_AxisLiteral`: a dotted tag in a symbol name gets re-parsed as an attribute access, leaving bare sympy symbols (no `dtype`) in an array's free symbols. * `IndexConnectorFmt` was mangled at one of its three uses, so a connector was declared under one name and referenced under another. `codegen_name` also escapes `[` and `]`: `Staggered[pkg.K]` has them and they survived into identifiers (Python read the emitted name as a subscript). The exhaustive injectivity test now covers `{a . _ [ ] u d l r}`. Tests: the DaCe tests that hard-coded symbol names -- 190 in one file -- now compute them from `sdfg_args`, so they stop depending on a naming scheme they are not testing. The binding golden test formats both sides, since the longer names make lines wrap. More dimensions that were value-equal across library modules are unified: `past_common` and `fvm_nabla_setup` import theirs, and `test_dace_bindings` takes the `cases` dimensions its fixture is sized on. --- .../next/0028-Dimensions_As_Nominal_Types.md | 8 +- src/gt4py/next/common.py | 10 +- src/gt4py/next/iterator/embedded.py | 3 +- .../codegens/gtfn/codegen.py | 3 +- .../codegens/gtfn/gtfn_module.py | 12 +- .../codegens/gtfn/itir_to_gtfn_ir.py | 4 +- .../dace/lowering/gtir_python_codegen.py | 6 +- .../dace/lowering/gtir_to_sdfg_lambda.py | 4 +- .../runners/dace/sdfg_args.py | 10 +- .../runners/dace/workflow/bindings.py | 6 +- tests/next_tests/fixtures/past_common.py | 9 +- .../ffront_tests/test_import_from_mod.py | 4 +- .../multi_feature_tests/fvm_nabla_setup.py | 16 +- .../dace_tests/test_dace_bindings.py | 62 ++- .../dace_tests/test_dace_translation.py | 24 +- .../dace_tests/test_gtir_to_sdfg.py | 470 ++++++++++-------- tests/next_tests/unit_tests/test_common.py | 15 +- 17 files changed, 385 insertions(+), 281 deletions(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 93311b6b38..299f4a76b5 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -116,17 +116,19 @@ disappears. ```python def codegen_name(tag: Tag) -> str: - return tag.replace("_", "_u").replace(".", "_d") + return tag.replace("_", "_u").replace(".", "_d").replace("[", "_l").replace("]", "_r") def from_codegen_name(name: str) -> Tag: - return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) + return re.sub(r"_([udlr])", lambda m: {"u": "_", "d": ".", "l": "[", "r": "]"}[m.group(1)], name) ``` A *prefix* escape, not `_ -> __` followed by `. -> _`: the latter is **not injective**, since a dot becomes a single underscore and `".."` collides with an escaped `"_"`. Every `_` in the output is the first character of a - two-character escape, so decoding is unambiguous. Names grow, which is what + two-character escape, so decoding is unambiguous. Brackets are escaped too: a + parametrized tag such as `Staggered[pkg.K]` contains them, and they would + otherwise survive into the identifier. Names grow, which is what gtfn's existing `TagDefinition.alias` mechanism is for. 6. **Staggered dimensions become a real parametrized type.** ADR 0026's diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index b64bdd28a4..983b702b67 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -63,6 +63,9 @@ class Dims(tuple[Unpack[ShapeTs]]): ... Tag: TypeAlias = str +_CODEGEN_UNESCAPE: Final = {"u": "_", "d": ".", "l": "[", "r": "]"} + + def codegen_name(tag: Tag) -> str: """ Mangle a dimension or offset tag into a valid generated identifier. @@ -90,7 +93,10 @@ def codegen_name(tag: Tag) -> str: >>> from_codegen_name(codegen_name("my__mod.X")) 'my__mod.X' """ - return tag.replace("_", "_u").replace(".", "_d") + # NOTE: `_` first, so the underscores introduced by the other escapes are not re-escaped. + # A tag's alphabet is `[A-Za-z0-9_.[]]`: brackets come from a parametrized tag such as + # `Staggered[pkg.K]`, and would otherwise survive into the identifier. + return tag.replace("_", "_u").replace(".", "_d").replace("[", "_l").replace("]", "_r") def from_codegen_name(name: str) -> Tag: @@ -110,7 +116,7 @@ def from_codegen_name(name: str) -> Tag: >>> from_codegen_name("mod_dV2E_dLocal") 'mod.V2E.Local' """ - return re.sub(r"_([ud])", lambda m: "_" if m.group(1) == "u" else ".", name) + return re.sub(r"_([udlr])", lambda m: _CODEGEN_UNESCAPE[m.group(1)], name) @enum.unique diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 6a21853f45..9b2179e243 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -1157,7 +1157,8 @@ def restrict(self, item: common.AnyIndexSpec) -> Self: if isinstance(item, Sequence) and all(isinstance(e, common.DimensionIndex) for e in item): assert len(item) == 1 assert isinstance(item[0], common.DimensionIndex) # for mypy errors on multiple lines below - d, r = item[0] + # an index is a `DimensionIndex` instance now, not a (dim, value) namedtuple + d, r = item[0].dim, item[0].value assert d == self._dimension assert isinstance(r, core_defs.INTEGRAL_TYPES) # TODO(tehrengruber): Use a regular zero dimensional field instead. diff --git a/src/gt4py/next/program_processors/codegens/gtfn/codegen.py b/src/gt4py/next/program_processors/codegens/gtfn/codegen.py index f44cffda4e..2647165b63 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/codegen.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/codegen.py @@ -117,7 +117,8 @@ def visit_Literal(self, node: gtfn_ir.Literal, **kwargs: Any) -> str: case "bool": return node.value.lower() case "axis_literal": - return node.value + # a qualified tag names a `generated::_t` tag type: mangle it as declared + return common.codegen_name(node.value) case _: # TODO(tehrengruber): we should probably shouldn't just allow anything here. Revisit. return node.value 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 5bc10c2fb5..965b376c94 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -110,10 +110,16 @@ def _process_connectivity_args( "Neighbor table indices must be of type 'np.int32' or 'np.int64'." ) + # NOTE: `name` is the offset-provider key, a qualified tag, so every identifier + # derived from it is mangled -- and identically, since the parameter name is + # referenced below and the `generated::_t` tag type is declared elsewhere. + cname = common.codegen_name(name) + param_name = GENERATED_CONNECTIVITY_PARAM_PREFIX + cname.lower() + # parameter parameters.append( interface.Parameter( - name=GENERATED_CONNECTIVITY_PARAM_PREFIX + name.lower(), + name=param_name, type_=ts.FieldType( dims=list(connectivity_type.domain), dtype=type_translation.from_dtype(connectivity_type.dtype), @@ -127,10 +133,10 @@ def _process_connectivity_args( f"generated::{common.codegen_name(connectivity_type.domain[0].tag)}_t, " f"generated::{common.codegen_name(connectivity_type.domain[1].tag)}_t, " f"{connectivity_type.max_neighbors}" - f">(std::forward({GENERATED_CONNECTIVITY_PARAM_PREFIX}{name.lower()}))" + f">(std::forward({param_name}))" ) arg_exprs.append( - f"gridtools::hymap::keys::make_values({nbtbl})" + f"gridtools::hymap::keys::make_values({nbtbl})" ) else: raise AssertionError( 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 0aec096134..7910632191 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 @@ -492,7 +492,9 @@ def _visit_unstructured_domain(self, node: itir.FunCall, **kwargs: Any) -> Node: common.get_offset_type(self.offset_provider_type, o), common.NeighborConnectivityType, ): - connectivities.append(SymRef(id=o)) + # `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. + connectivities.append(SymRef(id=common.codegen_name(o))) return UnstructuredDomain( tagged_sizes=sizes, tagged_offsets=domain_offsets, connectivities=connectivities ) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_python_codegen.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_python_codegen.py index 173d3a246c..6326351e23 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_python_codegen.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_python_codegen.py @@ -15,6 +15,7 @@ from gt4py.eve import codegen from gt4py.eve.codegen import FormatTemplate as as_fmt +from gt4py.next import common as gtx_common from gt4py.next.iterator import builtins, ir as gtir from gt4py.next.iterator.ir_utils import common_pattern_matcher as cpm @@ -135,7 +136,10 @@ class PythonCodegen(codegen.TemplatedGenerator): Literal = as_fmt("{value}") def visit_AxisLiteral(self, node: gtir.AxisLiteral, **kwargs: Any) -> str: - return node.value + # NOTE: mangled, because the result becomes part of a DaCe symbol name. A qualified tag + # there contains dots, which DaCe re-parses as attribute access, leaving plain sympy + # symbols (with no `dtype`) among an array's free symbols. + return gtx_common.codegen_name(node.value) def visit_FunCall(self, node: gtir.FunCall, args_map: dict[str, gtir.Node]) -> str: if cpm.is_let(node): 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 5e8a23d277..ed3ff892ec 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 @@ -665,7 +665,7 @@ def _visit_deref(self, node: gtir.FunCall) -> DataExpr: index_internals = ",".join( str(index.value - offset) if isinstance(index := arg_expr.indices[dim], SymbolExpr) - else f"{IndexConnectorFmt.format(dim=dim.tag)} - {offset}" + else f"{IndexConnectorFmt.format(dim=gtx_common.codegen_name(dim.tag))} - {offset}" for (dim, offset) in arg_expr.field_domain ) deref_node, connector_mapping = self._add_tasklet( @@ -684,7 +684,7 @@ def _visit_deref(self, node: gtir.FunCall) -> DataExpr: # add termination points for the dynamic iterator indices for dim, index_expr in field_indices: - index_connector = IndexConnectorFmt.format(dim=dim.tag) + index_connector = IndexConnectorFmt.format(dim=gtx_common.codegen_name(dim.tag)) if isinstance(index_expr, MemletExpr): self._add_input_data_edge( index_expr.dc_node, 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 4b3158c21b..5b0b0601fc 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -54,7 +54,9 @@ def as_itir_type(dtype: dace.typeclass) -> ts.ScalarType: def connectivity_identifier(name: str) -> str: - return f"{CONNECTIVITY_INDENTIFIER_PREFIX}{name}" + # NOTE: `name` is an offset-provider key, i.e. a qualified tag, which is not a valid SDFG + # array name; parse it back with `from_codegen_name` (see `is_connectivity_identifier`). + return f"{CONNECTIVITY_INDENTIFIER_PREFIX}{gtx_common.codegen_name(name)}" def is_connectivity_identifier( @@ -67,7 +69,7 @@ def is_connectivity_identifier( # that matches the CONNECTIVITY_INDENTIFIER_RE. return True else: - return gtx_common.has_offset(offset_provider_type, m[1]) + return gtx_common.has_offset(offset_provider_type, gtx_common.from_codegen_name(m[1])) def _field_symbol( @@ -80,8 +82,8 @@ def _field_symbol( name = f"__{field_name}_{gtx_common.codegen_name(dim.tag)}_{sym}" else: # a connectivity field assert offset_provider_type is not None - assert m[1] in offset_provider_type - offset = m[1] + 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: diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/bindings.py b/src/gt4py/next/program_processors/runners/dace/workflow/bindings.py index 82e6b3fc3c..77a1644b5d 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/bindings.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/bindings.py @@ -13,6 +13,7 @@ import dace from gt4py.eve import codegen +from gt4py.next import common as gtx_common from gt4py.next.otf import artifacts from gt4py.next.program_processors.runners.dace import ( sdfg_args as gtx_dace_args, @@ -206,8 +207,11 @@ def _parse_gt_connectivities( origin_size_param = next(iter(origin_size_arg.free_symbols)) m = gtx_dace_args.CONNECTIVITY_INDENTIFIER_RE.match(arg_name) assert m is not None + # NOTE: `m[1]` is the mangled offset name. It is a valid part of the Python variable + # name, but the offset provider is keyed by the real tag, so the lookup unmangles it. conn_arg = f"{_cb_neighbor_table}_{m[1]}" - code.append(f'{conn_arg} = {_cb_offset_provider}["{m[1]}"]') + offset = gtx_common.from_codegen_name(m[1]) + code.append(f'{conn_arg} = {_cb_offset_provider}["{offset}"]') _update_sdfg_array_ptr(code, conn_arg, sdfg_arg_index) _parse_gt_param( # set the size in the horizontal dimension param_name=origin_size_param, diff --git a/tests/next_tests/fixtures/past_common.py b/tests/next_tests/fixtures/past_common.py index 83c13988fe..43434f36ed 100644 --- a/tests/next_tests/fixtures/past_common.py +++ b/tests/next_tests/fixtures/past_common.py @@ -13,11 +13,10 @@ import gt4py.next as gtx from gt4py.next import float64 - -class IDim(gtx.DimensionIndex): ... - - -class JDim(gtx.DimensionIndex): ... +# NOTE: imported, not redeclared. Under nominal identity (ADR 0028) a same-named declaration +# here would be a different dimension from the one `cases_utils` declares, where the old +# `Dimension("...")` values compared equal -- and tests mix objects from both modules. +from next_tests.integration_tests.cases_utils import IDim # TODO(tehrengruber): Improve test structure. Identity needs to be decorated diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py index b41f3a9c89..262751f234 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py @@ -60,7 +60,7 @@ def test_import_offset_module_unstructured_shift(unstructured_case): def testee(a: cases.EField) -> cases.VField: return neighbor_sum(a(cases.V2E), axis=cases.V2EDim) - v2e_table = unstructured_case.offset_provider["V2E"].asnumpy() + v2e_table = unstructured_case.offset_provider[cases.V2EDim.tag].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -77,7 +77,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider["E2V"].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[cases.E2VDim.tag].asnumpy()[:, 0]], ) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py index 8c3a8095b8..4fbb01c72d 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py @@ -10,6 +10,7 @@ import numpy as np + try: from atlas4py import ( Config, @@ -35,17 +36,10 @@ from gt4py import next as gtx from gt4py.next.iterator import atlas_utils - -class Vertex(gtx.DimensionIndex): ... - - -class Edge(gtx.DimensionIndex): ... - - -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... - - -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +# NOTE: imported, not redeclared. Under nominal identity (ADR 0028) a same-named declaration +# here would be a different dimension from the one `toy_connectivity` declares, where the old +# `Dimension("...")` values compared equal -- and tests mix objects from both modules. +from next_tests.toy_connectivity import E2VDim, Edge, V2EDim, Vertex V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index 922ffbf3ff..6ecaa45fcc 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -9,22 +9,23 @@ """Test the bindings stage of the dace backend workflow.""" import functools -import dace + import numpy as np import pytest -from gt4py.eve import codegen from gt4py import next as gtx -from gt4py.next import common as gtx_common, int32 +from gt4py.eve import codegen +from gt4py.next import common as gtx_common, int32, neighbor_sum from gt4py.next.otf import artifacts from gt4py.next.program_processors.runners import dace as dace_runner from gt4py.next.program_processors.runners.dace import workflow as dace_workflow -from gt4py.next import neighbor_sum -from next_tests.integration_tests.cases import E2V, E2VDim, V2E, V2EDim -from next_tests.integration_tests import cases -from next_tests.integration_tests import cases_utils -from next_tests.unit_tests.test_common import IDim, JDim, KDim +from next_tests.integration_tests import cases, cases_utils + +# NOTE: from `cases`, not `test_common`: the `cartesian_case` fixture is sized on the `cases` +# dimensions, and under nominal identity (ADR 0028) another module's same-named `IDim` is a +# different dimension -- it used to compare equal. +from next_tests.integration_tests.cases import E2V, V2E, E2VDim, IDim, JDim, KDim, V2EDim _bind_func_name = "update_sdfg_args" @@ -154,6 +155,12 @@ def {_bind_func_name}(device, sdfg_argtypes, args, sdfg_call_args, offset_provid ) +# The generated binding names each table variable after the *mangled* offset key -- it has to be +# a valid identifier -- but looks the table up in the offset provider by the real, qualified tag. +_E2V_TABLE = f"table_{gtx_common.codegen_name(E2VDim.tag)}" +_V2E_TABLE = f"table_{gtx_common.codegen_name(V2EDim.tag)}" + + def _binding_source_unstructured(use_metrics: bool) -> str: metrics_arg_index = 2 idx = [0, 4, 5, 1, 6, 7, 8, 2, 10, 9, 3, 12, 11] @@ -174,14 +181,14 @@ def {_bind_func_name}(device, sdfg_argtypes, args, sdfg_call_args, offset_provid sdfg_call_args[{idx[4]}] = ctypes.c_int(args_1.domain.ranges[0].start) sdfg_call_args[{idx[5]}] = ctypes.c_int(args_1.domain.ranges[0].stop) sdfg_call_args[{idx[6]}] = ctypes.c_int(args_1.__gt_buffer_info__.elem_strides[0]) - table_E2V = offset_provider["E2V"] - sdfg_call_args[{idx[7]}].value = table_E2V.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[8]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[0]) - sdfg_call_args[{idx[9]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[1]) - table_V2E = offset_provider["V2E"] - sdfg_call_args[{idx[10]}].value = table_V2E.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[11]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[0]) - sdfg_call_args[{idx[12]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[1]) + {_E2V_TABLE} = offset_provider["{E2VDim.tag}"] + sdfg_call_args[{idx[7]}].value = {_E2V_TABLE}.__gt_buffer_info__.data_ptr + sdfg_call_args[{idx[8]}] = ctypes.c_int({_E2V_TABLE}.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[9]}] = ctypes.c_int({_E2V_TABLE}.__gt_buffer_info__.elem_strides[1]) + {_V2E_TABLE} = offset_provider["{V2EDim.tag}"] + sdfg_call_args[{idx[10]}].value = {_V2E_TABLE}.__gt_buffer_info__.data_ptr + sdfg_call_args[{idx[11]}] = ctypes.c_int({_V2E_TABLE}.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[12]}] = ctypes.c_int({_V2E_TABLE}.__gt_buffer_info__.elem_strides[1]) """ ) @@ -204,14 +211,14 @@ def {_bind_func_name}(device, sdfg_argtypes, args, sdfg_call_args, offset_provid sdfg_call_args[{idx[2]}].value = args_1.__gt_buffer_info__.data_ptr sdfg_call_args[{idx[3]}] = ctypes.c_int(args_1.domain.ranges[0].stop) sdfg_call_args[{idx[4]}] = ctypes.c_int(args_1.__gt_buffer_info__.elem_strides[0]) - table_E2V = offset_provider["E2V"] - sdfg_call_args[{idx[5]}].value = table_E2V.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[6]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[0]) - sdfg_call_args[{idx[7]}] = ctypes.c_int(table_E2V.__gt_buffer_info__.elem_strides[1]) - table_V2E = offset_provider["V2E"] - sdfg_call_args[{idx[8]}].value = table_V2E.__gt_buffer_info__.data_ptr - sdfg_call_args[{idx[9]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[0]) - sdfg_call_args[{idx[10]}] = ctypes.c_int(table_V2E.__gt_buffer_info__.elem_strides[1]) + {_E2V_TABLE} = offset_provider["{E2VDim.tag}"] + sdfg_call_args[{idx[5]}].value = {_E2V_TABLE}.__gt_buffer_info__.data_ptr + sdfg_call_args[{idx[6]}] = ctypes.c_int({_E2V_TABLE}.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[7]}] = ctypes.c_int({_E2V_TABLE}.__gt_buffer_info__.elem_strides[1]) + {_V2E_TABLE} = offset_provider["{V2EDim.tag}"] + sdfg_call_args[{idx[8]}].value = {_V2E_TABLE}.__gt_buffer_info__.data_ptr + sdfg_call_args[{idx[9]}] = ctypes.c_int({_V2E_TABLE}.__gt_buffer_info__.elem_strides[0]) + sdfg_call_args[{idx[10]}] = ctypes.c_int({_V2E_TABLE}.__gt_buffer_info__.elem_strides[1]) """ ) @@ -238,7 +245,12 @@ def mocked_compile_call( for line in inp.binding_source.source_code.splitlines() if not line.lstrip().startswith("assert") ) - assert codegen.format_python_source(binding_source_pruned) == binding_source_ref + # NOTE: both sides go through the same formatter. The generated side always did; formatting + # the reference too makes the comparison independent of where long lines happen to wrap, + # which the mangled (qualified) connectivity names in the binding now make them do. + assert codegen.format_python_source(binding_source_pruned) == codegen.format_python_source( + binding_source_ref + ) return _dace_compile_call(self, inp) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index 87e212afc9..1bec200fad 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -8,23 +8,24 @@ """Test the translation stage of the dace backend workflow.""" -import dace -import pytest - import re import uuid from typing import Callable from unittest import mock +import dace +import pytest +from dace import nodes as dace_nodes + from gt4py._core import definitions as core_defs from gt4py.next import common as gtx_common, fingerprinting from gt4py.next.iterator import ir as itir from gt4py.next.iterator.ir_utils import ir_makers as im from gt4py.next.otf import arguments as otf_arguments, workflow as otf_workflow -from gt4py.next.program_processors.runners.dace import lowering as gtx_dace_lowering +from gt4py.next.program_processors.runners.dace import sdfg_args as gtx_dace_args from gt4py.next.program_processors.runners.dace.workflow import ( - translation as dace_wf_translation, common as dace_wf_common, + translation as dace_wf_translation, ) from gt4py.next.type_system import type_specifications as ts @@ -32,12 +33,11 @@ V2E, Edge, IDim, + V2EDim, Vertex, skip_value_mesh, ) -from dace import nodes as dace_nodes - FLOAT_TYPE = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) IFTYPE = ts.FieldType(dims=[IDim], dtype=FLOAT_TYPE) @@ -120,14 +120,14 @@ def test_find_constant_symbols(has_unit_stride, disable_field_origin): expected = {} if has_unit_stride: expected |= { - "__x_Edge_stride": 1, - "__y_Vertex_stride": 1, - "__gt_conn_V2E_source_stride": 1, + gtx_dace_args.field_stride_symbol("x", Edge).name: 1, + gtx_dace_args.field_stride_symbol("y", Vertex).name: 1, + f"__{gtx_dace_args.connectivity_identifier(V2EDim.tag)}_source_stride": 1, } if disable_field_origin: expected |= { - "__x_Edge_range_0": 0, - "__y_Vertex_range_0": 0, + gtx_dace_args.range_start_symbol("x", Edge).name: 0, + gtx_dace_args.range_start_symbol("y", Vertex).name: 0, } assert constant_symbols == expected diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py index 8e9d037aa6..a1fb215882 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py @@ -43,6 +43,7 @@ ) from gt4py.next.program_processors.runners.dace import lowering as dace_lowering +from gt4py.next.program_processors.runners.dace import sdfg_args as gtx_dace_args @pytest.fixture @@ -75,18 +76,18 @@ def allow_view_arguments(): SKIP_VALUE_MESH: MeshDescriptor = skip_value_mesh(None) SIZE_TYPE = ts.ScalarType(ts.ScalarKind.INT32) FSYMBOLS = dict( - __w_IDim_range_0=0, - __w_IDim_range_1=N, - __w_IDim_stride=1, - __x_IDim_range_0=0, - __x_IDim_range_1=N, - __x_IDim_stride=1, - __y_IDim_range_0=0, - __y_IDim_range_1=N, - __y_IDim_stride=1, - __z_IDim_range_0=0, - __z_IDim_range_1=N, - __z_IDim_stride=1, + **{gtx_dace_args.range_start_symbol("w", IDim).name: 0}, + **{gtx_dace_args.range_stop_symbol("w", IDim).name: N}, + **{gtx_dace_args.field_stride_symbol("w", IDim).name: 1}, + **{gtx_dace_args.range_start_symbol("x", IDim).name: 0}, + **{gtx_dace_args.range_stop_symbol("x", IDim).name: N}, + **{gtx_dace_args.field_stride_symbol("x", IDim).name: 1}, + **{gtx_dace_args.range_start_symbol("y", IDim).name: 0}, + **{gtx_dace_args.range_stop_symbol("y", IDim).name: N}, + **{gtx_dace_args.field_stride_symbol("y", IDim).name: 1}, + **{gtx_dace_args.range_start_symbol("z", IDim).name: 0}, + **{gtx_dace_args.range_stop_symbol("z", IDim).name: N}, + **{gtx_dace_args.field_stride_symbol("z", IDim).name: 1}, ) @@ -96,27 +97,83 @@ def make_mesh_symbols(mesh: MeshDescriptor): e2v_ndarray = mesh.offset_provider[E2VDim.tag].ndarray v2e_ndarray = mesh.offset_provider[V2EDim.tag].ndarray return dict( - __cells_Cell_range_0=0, - __cells_Cell_range_1=mesh.num_cells, - __cells_Cell_stride=1, - __edges_Edge_range_0=0, - __edges_Edge_range_1=mesh.num_edges, - __edges_Edge_stride=1, - __vertices_Vertex_range_0=0, - __vertices_Vertex_range_1=mesh.num_vertices, - __vertices_Vertex_stride=1, - __gt_conn_C2E_source_size=c2e_ndarray.shape[0], - __gt_conn_C2E_source_stride=c2e_ndarray.strides[0] // c2e_ndarray.itemsize, - __gt_conn_C2E_neighbor_stride=c2e_ndarray.strides[1] // c2e_ndarray.itemsize, - __gt_conn_C2V_source_size=c2v_ndarray.shape[0], - __gt_conn_C2V_source_stride=c2v_ndarray.strides[0] // c2v_ndarray.itemsize, - __gt_conn_C2V_neighbor_stride=c2v_ndarray.strides[1] // c2v_ndarray.itemsize, - __gt_conn_E2V_source_size=e2v_ndarray.shape[0], - __gt_conn_E2V_source_stride=e2v_ndarray.strides[0] // e2v_ndarray.itemsize, - __gt_conn_E2V_neighbor_stride=e2v_ndarray.strides[1] // e2v_ndarray.itemsize, - __gt_conn_V2E_source_size=v2e_ndarray.shape[0], - __gt_conn_V2E_source_stride=v2e_ndarray.strides[0] // v2e_ndarray.itemsize, - __gt_conn_V2E_neighbor_stride=v2e_ndarray.strides[1] // v2e_ndarray.itemsize, + **{gtx_dace_args.range_start_symbol("cells", Cell).name: 0}, + **{gtx_dace_args.range_stop_symbol("cells", Cell).name: mesh.num_cells}, + **{gtx_dace_args.field_stride_symbol("cells", Cell).name: 1}, + **{gtx_dace_args.range_start_symbol("edges", Edge).name: 0}, + **{gtx_dace_args.range_stop_symbol("edges", Edge).name: mesh.num_edges}, + **{gtx_dace_args.field_stride_symbol("edges", Edge).name: 1}, + **{gtx_dace_args.range_start_symbol("vertices", Vertex).name: 0}, + **{gtx_dace_args.range_stop_symbol("vertices", Vertex).name: mesh.num_vertices}, + **{gtx_dace_args.field_stride_symbol("vertices", Vertex).name: 1}, + **{ + f"__{gtx_dace_args.connectivity_identifier(C2EDim.tag)}_source_size": c2e_ndarray.shape[ + 0 + ] + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(C2EDim.tag)}_source_stride": c2e_ndarray.strides[ + 0 + ] + // c2e_ndarray.itemsize + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(C2EDim.tag)}_neighbor_stride": c2e_ndarray.strides[ + 1 + ] + // c2e_ndarray.itemsize + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(C2VDim.tag)}_source_size": c2v_ndarray.shape[ + 0 + ] + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(C2VDim.tag)}_source_stride": c2v_ndarray.strides[ + 0 + ] + // c2v_ndarray.itemsize + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(C2VDim.tag)}_neighbor_stride": c2v_ndarray.strides[ + 1 + ] + // c2v_ndarray.itemsize + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(E2VDim.tag)}_source_size": e2v_ndarray.shape[ + 0 + ] + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(E2VDim.tag)}_source_stride": e2v_ndarray.strides[ + 0 + ] + // e2v_ndarray.itemsize + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(E2VDim.tag)}_neighbor_stride": e2v_ndarray.strides[ + 1 + ] + // e2v_ndarray.itemsize + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(V2EDim.tag)}_source_size": v2e_ndarray.shape[ + 0 + ] + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(V2EDim.tag)}_source_stride": v2e_ndarray.strides[ + 0 + ] + // v2e_ndarray.itemsize + }, + **{ + f"__{gtx_dace_args.connectivity_identifier(V2EDim.tag)}_neighbor_stride": v2e_ndarray.strides[ + 1 + ] + // v2e_ndarray.itemsize + }, ) @@ -311,15 +368,15 @@ def test_gtir_tuple_args(): x_fields = (a, a, b) tuple_symbols = { - "__x_0_IDim_range_0": 0, - "__x_0_IDim_range_1": N, - "__x_0_IDim_stride": 1, - "__x_1_0_IDim_range_0": 0, - "__x_1_0_IDim_range_1": N, - "__x_1_0_IDim_stride": 1, - "__x_1_1_IDim_range_0": 0, - "__x_1_1_IDim_range_1": N, - "__x_1_1_IDim_stride": 1, + gtx_dace_args.range_start_symbol("x_0", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_0", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("x_1_0", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_1_0", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_1_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("x_1_1", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_1_1", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_1_1", IDim).name: 1, } sdfg(*x_fields, c, **FSYMBOLS, **tuple_symbols) @@ -496,15 +553,15 @@ def test_gtir_tuple_return(): z_fields = (np.empty_like(a), np.empty_like(a), np.empty_like(a)) tuple_symbols = { - "__z_0_0_IDim_range_0": 0, - "__z_0_0_IDim_range_1": N, - "__z_0_0_IDim_stride": 1, - "__z_0_1_IDim_range_0": 0, - "__z_0_1_IDim_range_1": N, - "__z_0_1_IDim_stride": 1, - "__z_1_IDim_range_0": 0, - "__z_1_IDim_range_1": N, - "__z_1_IDim_stride": 1, + gtx_dace_args.range_start_symbol("z_0_0", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_0_0", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_0_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("z_0_1", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_0_1", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_0_1", IDim).name: 1, + gtx_dace_args.range_start_symbol("z_1", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_1", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_1", IDim).name: 1, } sdfg(a, b, *z_fields, **FSYMBOLS, **tuple_symbols) @@ -758,12 +815,12 @@ def test_gtir_cond_with_tuple_return(): sdfg = build_dace_sdfg(testee, CARTESIAN_OFFSETS) tuple_symbols = { - "__z_0_IDim_range_0": 0, - "__z_0_IDim_range_1": N, - "__z_0_IDim_stride": 1, - "__z_1_IDim_range_0": 0, - "__z_1_IDim_range_1": N, - "__z_1_IDim_stride": 1, + gtx_dace_args.range_start_symbol("z_0", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_0", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("z_1", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_1", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_1", IDim).name: 1, } for s in [False, True]: @@ -894,9 +951,9 @@ def test_gtir_cartesian_shift_left(): sdfg = build_dace_sdfg(testee, CARTESIAN_OFFSETS) symbols = FSYMBOLS | { - "__x_offset_IDim_range_0": 0, - "__x_offset_IDim_range_1": N, - "__x_offset_IDim_stride": 1, + gtx_dace_args.range_start_symbol("x_offset", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_offset", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_offset", IDim).name: 1, } sdfg(a, a_offset, b, **symbols) @@ -985,9 +1042,9 @@ def test_gtir_cartesian_shift_right(): sdfg = build_dace_sdfg(testee, CARTESIAN_OFFSETS) symbols = FSYMBOLS | { - "__x_offset_IDim_range_0": 0, - "__x_offset_IDim_range_1": N, - "__x_offset_IDim_stride": 1, + gtx_dace_args.range_start_symbol("x_offset", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_offset", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_offset", IDim).name: 1, } sdfg(a, a_offset, b, **symbols) @@ -1116,28 +1173,28 @@ def test_gtir_connectivity_shift(): ev, c2e_offset=np.full(SIMPLE_MESH.num_cells, C2E_neighbor_idx, dtype=np.int32), e2v_offset=np.full(SIMPLE_MESH.num_edges, E2V_neighbor_idx, dtype=np.int32), - gt_conn_C2E=connectivity_C2E.ndarray, - gt_conn_E2V=connectivity_E2V.ndarray, + **{gtx_dace_args.connectivity_identifier(C2EDim.tag): connectivity_C2E.ndarray}, + **{gtx_dace_args.connectivity_identifier(E2VDim.tag): connectivity_E2V.ndarray}, **FSYMBOLS, **make_mesh_symbols(SIMPLE_MESH), - __ce_field_Cell_range_0=0, - __ce_field_Cell_range_1=SIMPLE_MESH.num_cells, - __ce_field_Cell_stride=SIMPLE_MESH.num_edges, - __ce_field_Edge_range_0=0, - __ce_field_Edge_range_1=SIMPLE_MESH.num_edges, - __ce_field_Edge_stride=1, - __ev_field_Edge_range_0=0, - __ev_field_Edge_range_1=SIMPLE_MESH.num_edges, - __ev_field_Edge_stride=SIMPLE_MESH.num_vertices, - __ev_field_Vertex_range_0=0, - __ev_field_Vertex_range_1=SIMPLE_MESH.num_vertices, - __ev_field_Vertex_stride=1, - __c2e_offset_Cell_range_0=0, - __c2e_offset_Cell_range_1=SIMPLE_MESH.num_cells, - __c2e_offset_Cell_stride=1, - __e2v_offset_Edge_range_0=0, - __e2v_offset_Edge_range_1=SIMPLE_MESH.num_edges, - __e2v_offset_Edge_stride=1, + **{gtx_dace_args.range_start_symbol("ce_field", Cell).name: 0}, + **{gtx_dace_args.range_stop_symbol("ce_field", Cell).name: SIMPLE_MESH.num_cells}, + **{gtx_dace_args.field_stride_symbol("ce_field", Cell).name: SIMPLE_MESH.num_edges}, + **{gtx_dace_args.range_start_symbol("ce_field", Edge).name: 0}, + **{gtx_dace_args.range_stop_symbol("ce_field", Edge).name: SIMPLE_MESH.num_edges}, + **{gtx_dace_args.field_stride_symbol("ce_field", Edge).name: 1}, + **{gtx_dace_args.range_start_symbol("ev_field", Edge).name: 0}, + **{gtx_dace_args.range_stop_symbol("ev_field", Edge).name: SIMPLE_MESH.num_edges}, + **{gtx_dace_args.field_stride_symbol("ev_field", Edge).name: SIMPLE_MESH.num_vertices}, + **{gtx_dace_args.range_start_symbol("ev_field", Vertex).name: 0}, + **{gtx_dace_args.range_stop_symbol("ev_field", Vertex).name: SIMPLE_MESH.num_vertices}, + **{gtx_dace_args.field_stride_symbol("ev_field", Vertex).name: 1}, + **{gtx_dace_args.range_start_symbol("c2e_offset", Cell).name: 0}, + **{gtx_dace_args.range_stop_symbol("c2e_offset", Cell).name: SIMPLE_MESH.num_cells}, + **{gtx_dace_args.field_stride_symbol("c2e_offset", Cell).name: 1}, + **{gtx_dace_args.range_start_symbol("e2v_offset", Edge).name: 0}, + **{gtx_dace_args.range_stop_symbol("e2v_offset", Edge).name: SIMPLE_MESH.num_edges}, + **{gtx_dace_args.field_stride_symbol("e2v_offset", Edge).name: 1}, ) assert np.allclose(ce, ref) @@ -1193,13 +1250,13 @@ def test_gtir_connectivity_shift_chain(): sdfg( e, e_out, - gt_conn_E2V=connectivity_E2V.ndarray, - gt_conn_V2E=connectivity_V2E.ndarray, + **{gtx_dace_args.connectivity_identifier(E2VDim.tag): connectivity_E2V.ndarray}, + **{gtx_dace_args.connectivity_identifier(V2EDim.tag): connectivity_V2E.ndarray}, **FSYMBOLS, **make_mesh_symbols(SIMPLE_MESH), - __edges_out_Edge_range_0=0, - __edges_out_Edge_range_1=SIMPLE_MESH.num_edges, - __edges_out_Edge_stride=1, + **{gtx_dace_args.range_start_symbol("edges_out", Edge).name: 0}, + **{gtx_dace_args.range_stop_symbol("edges_out", Edge).name: SIMPLE_MESH.num_edges}, + **{gtx_dace_args.field_stride_symbol("edges_out", Edge).name: 1}, ) assert np.allclose(e_out, ref) @@ -1273,26 +1330,35 @@ def test_gtir_neighbors_as_input(): symbols = make_mesh_symbols(SIMPLE_MESH) | { # override SDFG symbols for array shape and strides because of extra K-dimension - "__edges_KDim_range_0": 0, - "__edges_KDim_range_1": e.shape[1], - "__edges_Edge_stride": e.strides[0] // e.itemsize, - "__edges_KDim_stride": e.strides[1] // e.itemsize, - "__vertices_KDim_range_0": 0, - "__vertices_KDim_range_1": v.shape[1], - "__vertices_Vertex_stride": v.strides[0] // v.itemsize, - "__vertices_KDim_stride": v.strides[1] // v.itemsize, - "__v2e_field_Vertex_range_0": 0, - "__v2e_field_Vertex_range_1": v2e_field.shape[0], - "__v2e_field_Vertex_stride": v2e_field.strides[0] // v2e_field.itemsize, - "__v2e_field_V2E_range_0": 0, - "__v2e_field_V2E_range_1": v2e_field.shape[1], - "__v2e_field_V2E_stride": v2e_field.strides[1] // v2e_field.itemsize, - "__v2e_field_KDim_range_0": 0, - "__v2e_field_KDim_range_1": v2e_field.shape[2], - "__v2e_field_KDim_stride": v2e_field.strides[2] // v2e_field.itemsize, + gtx_dace_args.range_start_symbol("edges", KDim).name: 0, + gtx_dace_args.range_stop_symbol("edges", KDim).name: e.shape[1], + gtx_dace_args.field_stride_symbol("edges", Edge).name: e.strides[0] // e.itemsize, + gtx_dace_args.field_stride_symbol("edges", KDim).name: e.strides[1] // e.itemsize, + gtx_dace_args.range_start_symbol("vertices", KDim).name: 0, + gtx_dace_args.range_stop_symbol("vertices", KDim).name: v.shape[1], + gtx_dace_args.field_stride_symbol("vertices", Vertex).name: v.strides[0] // v.itemsize, + gtx_dace_args.field_stride_symbol("vertices", KDim).name: v.strides[1] // v.itemsize, + gtx_dace_args.range_start_symbol("v2e_field", Vertex).name: 0, + gtx_dace_args.range_stop_symbol("v2e_field", Vertex).name: v2e_field.shape[0], + gtx_dace_args.field_stride_symbol("v2e_field", Vertex).name: v2e_field.strides[0] + // v2e_field.itemsize, + gtx_dace_args.range_start_symbol("v2e_field", V2EDim).name: 0, + gtx_dace_args.range_stop_symbol("v2e_field", V2EDim).name: v2e_field.shape[1], + gtx_dace_args.field_stride_symbol("v2e_field", V2EDim).name: v2e_field.strides[1] + // v2e_field.itemsize, + gtx_dace_args.range_start_symbol("v2e_field", KDim).name: 0, + gtx_dace_args.range_stop_symbol("v2e_field", KDim).name: v2e_field.shape[2], + gtx_dace_args.field_stride_symbol("v2e_field", KDim).name: v2e_field.strides[2] + // v2e_field.itemsize, } - sdfg(v2e_field, e, v, gt_conn_V2E=connectivity_V2E.ndarray, **symbols) + sdfg( + v2e_field, + e, + v, + **{gtx_dace_args.connectivity_identifier(V2EDim.tag): connectivity_V2E.ndarray}, + **symbols, + ) assert np.allclose(v, v_ref) @@ -1344,7 +1410,7 @@ def test_gtir_reduce(): sdfg( e, v, - gt_conn_V2E=connectivity_V2E.ndarray, + **{gtx_dace_args.connectivity_identifier(V2EDim.tag): connectivity_V2E.ndarray}, **FSYMBOLS, **make_mesh_symbols(SIMPLE_MESH), ) @@ -1401,7 +1467,7 @@ def test_gtir_reduce_with_skip_values(): sdfg( e, v, - gt_conn_V2E=connectivity_V2E.ndarray, + **{gtx_dace_args.connectivity_identifier(V2EDim.tag): connectivity_V2E.ndarray}, **FSYMBOLS, **make_mesh_symbols(SKIP_VALUE_MESH), ) @@ -1464,12 +1530,12 @@ def test_gtir_reduce_dot_product(): v2e_field, e, v, - gt_conn_V2E=connectivity_V2E.ndarray, + **{gtx_dace_args.connectivity_identifier(V2EDim.tag): connectivity_V2E.ndarray}, **make_mesh_symbols(SKIP_VALUE_MESH), - __v2e_field_Vertex_range_0=0, - __v2e_field_Vertex_range_1=SKIP_VALUE_MESH.num_vertices, - __v2e_field_Vertex_stride=connectivity_V2E.shape[1], - __v2e_field_V2E_stride=1, + **{gtx_dace_args.range_start_symbol("v2e_field", Vertex).name: 0}, + **{gtx_dace_args.range_stop_symbol("v2e_field", Vertex).name: SKIP_VALUE_MESH.num_vertices}, + **{gtx_dace_args.field_stride_symbol("v2e_field", Vertex).name: connectivity_V2E.shape[1]}, + **{gtx_dace_args.field_stride_symbol("v2e_field", V2EDim).name: 1}, ) assert np.allclose(v, v_ref) @@ -1535,13 +1601,13 @@ def test_gtir_reduce_with_cond_neighbors(use_sparse): v2e_field, e, v, - gt_conn_V2E=connectivity_V2E.ndarray, + **{gtx_dace_args.connectivity_identifier(V2EDim.tag): connectivity_V2E.ndarray}, **FSYMBOLS, **make_mesh_symbols(SKIP_VALUE_MESH), - __v2e_field_Vertex_range_0=0, - __v2e_field_Vertex_range_1=SKIP_VALUE_MESH.num_vertices, - __v2e_field_Vertex_stride=connectivity_V2E.shape[1], - __v2e_field_V2E_stride=1, + **{gtx_dace_args.range_start_symbol("v2e_field", Vertex).name: 0}, + **{gtx_dace_args.range_stop_symbol("v2e_field", Vertex).name: SKIP_VALUE_MESH.num_vertices}, + **{gtx_dace_args.field_stride_symbol("v2e_field", Vertex).name: connectivity_V2E.shape[1]}, + **{gtx_dace_args.field_stride_symbol("v2e_field", V2EDim).name: 1}, ) assert np.allclose(v, v_ref) @@ -1741,7 +1807,7 @@ def test_gtir_let_lambda_scalar_expression(): # to the symbol `inner_size` is preserved, for which we want to test the lowering. sdfg = build_dace_sdfg(testee, offset_provider=CARTESIAN_OFFSETS, skip_domain_inference=True) - sdfg(a, b, c, d, **(FSYMBOLS | {"__x_IDim_range_1": N + 1})) + sdfg(a, b, c, d, **(FSYMBOLS | {gtx_dace_args.range_stop_symbol("x", IDim).name: N + 1})) assert np.allclose(d, (a * a * b * b * c[1 : N + 1])) @@ -1796,8 +1862,8 @@ def test_gtir_let_lambda_with_connectivity(): cells=c, edges=e, vertices=v, - gt_conn_C2E=connectivity_C2E.ndarray, - gt_conn_C2V=connectivity_C2V.ndarray, + **{gtx_dace_args.connectivity_identifier(C2EDim.tag): connectivity_C2E.ndarray}, + **{gtx_dace_args.connectivity_identifier(C2VDim.tag): connectivity_C2V.ndarray}, **FSYMBOLS, **make_mesh_symbols(SIMPLE_MESH), ) @@ -1846,20 +1912,20 @@ def test_gtir_let_lambda_with_origin(): ) symbols = make_mesh_symbols(SIMPLE_MESH) | { - "__cells_KDim_range_0": 0, - "__cells_KDim_range_1": MESH_NUM_LEVELS, - "__cells_Cell_stride": c.strides[0] // c.itemsize, - "__cells_KDim_stride": c.strides[1] // c.itemsize, - "__edges_KDim_range_0": 0, - "__edges_KDim_range_1": MESH_NUM_LEVELS, - "__edges_Edge_stride": e.strides[0] // e.itemsize, - "__edges_KDim_stride": e.strides[1] // e.itemsize, + gtx_dace_args.range_start_symbol("cells", KDim).name: 0, + gtx_dace_args.range_stop_symbol("cells", KDim).name: MESH_NUM_LEVELS, + gtx_dace_args.field_stride_symbol("cells", Cell).name: c.strides[0] // c.itemsize, + gtx_dace_args.field_stride_symbol("cells", KDim).name: c.strides[1] // c.itemsize, + gtx_dace_args.range_start_symbol("edges", KDim).name: 0, + gtx_dace_args.range_stop_symbol("edges", KDim).name: MESH_NUM_LEVELS, + gtx_dace_args.field_stride_symbol("edges", Edge).name: e.strides[0] // e.itemsize, + gtx_dace_args.field_stride_symbol("edges", KDim).name: e.strides[1] // e.itemsize, } sdfg( cells=c, edges=e, - gt_conn_C2E=connectivity_C2E.ndarray, + **{gtx_dace_args.connectivity_identifier(C2EDim.tag): connectivity_C2E.ndarray}, **symbols, ) @@ -1942,12 +2008,12 @@ def test_gtir_let_lambda_with_tuple1(): b_ref = np.concatenate((z_fields[1][:1], b[1 : N - 1], z_fields[1][N - 1 :])) tuple_symbols = { - "__z_0_IDim_range_0": 1, - "__z_0_IDim_range_1": N - 1, - "__z_0_IDim_stride": 1, - "__z_1_IDim_range_0": 1, - "__z_1_IDim_range_1": N - 1, - "__z_1_IDim_stride": 1, + gtx_dace_args.range_start_symbol("z_0", IDim).name: 1, + gtx_dace_args.range_stop_symbol("z_0", IDim).name: N - 1, + gtx_dace_args.field_stride_symbol("z_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("z_1", IDim).name: 1, + gtx_dace_args.range_stop_symbol("z_1", IDim).name: N - 1, + gtx_dace_args.field_stride_symbol("z_1", IDim).name: 1, } sdfg(a, b, z_fields[0][1 : N - 1], z_fields[1][1 : N - 1], **FSYMBOLS, **tuple_symbols) @@ -2000,15 +2066,15 @@ def test_gtir_let_lambda_with_tuple2(): z_fields = (np.empty_like(a), np.empty_like(a), np.empty_like(a)) tuple_symbols = { - "__z_0_IDim_range_0": 0, - "__z_0_IDim_range_1": N, - "__z_0_IDim_stride": 1, - "__z_1_IDim_range_0": 0, - "__z_1_IDim_range_1": N, - "__z_1_IDim_stride": 1, - "__z_2_IDim_range_0": 0, - "__z_2_IDim_range_1": N, - "__z_2_IDim_stride": 1, + gtx_dace_args.range_start_symbol("z_0", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_0", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("z_1", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_1", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_1", IDim).name: 1, + gtx_dace_args.range_start_symbol("z_2", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z_2", IDim).name: N, + gtx_dace_args.field_stride_symbol("z_2", IDim).name: 1, } sdfg(a, b, *z_fields, **FSYMBOLS, **tuple_symbols) @@ -2063,15 +2129,15 @@ def test_gtir_if_scalars(s): sdfg = build_dace_sdfg(testee, {}) tuple_symbols = { - "__x_0_IDim_range_0": 0, - "__x_0_IDim_range_1": N, - "__x_0_IDim_stride": 1, - "__x_1_0_IDim_range_0": 0, - "__x_1_0_IDim_range_1": N, - "__x_1_0_IDim_stride": 1, - "__x_1_1_IDim_range_0": 0, - "__x_1_1_IDim_range_1": N, - "__x_1_1_IDim_stride": 1, + gtx_dace_args.range_start_symbol("x_0", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_0", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("x_1_0", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_1_0", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_1_0", IDim).name: 1, + gtx_dace_args.range_start_symbol("x_1_1", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x_1_1", IDim).name: N, + gtx_dace_args.field_stride_symbol("x_1_1", IDim).name: 1, } sdfg(x_0=a, x_1_0=d1, x_1_1=d2, z=b, pred=np.bool_(s), **FSYMBOLS, **tuple_symbols) @@ -2245,30 +2311,30 @@ def test_gtir_concat_where_two_dimensions(): ) field_symbols = { - "__x_IDim_range_0": 0, - "__x_IDim_range_1": a.shape[0], - "__x_JDim_range_0": 0, - "__x_JDim_range_1": a.shape[1], - "__x_IDim_stride": a.strides[0] // a.itemsize, - "__x_JDim_stride": a.strides[1] // a.itemsize, - "__y_IDim_range_0": 0, - "__y_IDim_range_1": b.shape[0], - "__y_JDim_range_0": 0, - "__y_JDim_range_1": b.shape[1], - "__y_IDim_stride": b.strides[0] // b.itemsize, - "__y_JDim_stride": b.strides[1] // b.itemsize, - "__w_IDim_range_0": 0, - "__w_IDim_range_1": c.shape[0], - "__w_JDim_range_0": 0, - "__w_JDim_range_1": c.shape[1], - "__w_IDim_stride": c.strides[0] // c.itemsize, - "__w_JDim_stride": c.strides[1] // c.itemsize, - "__z_IDim_range_0": 0, - "__z_IDim_range_1": d.shape[0], - "__z_JDim_range_0": 0, - "__z_JDim_range_1": d.shape[1], - "__z_IDim_stride": d.strides[0] // d.itemsize, - "__z_JDim_stride": d.strides[1] // d.itemsize, + gtx_dace_args.range_start_symbol("x", IDim).name: 0, + gtx_dace_args.range_stop_symbol("x", IDim).name: a.shape[0], + gtx_dace_args.range_start_symbol("x", JDim).name: 0, + gtx_dace_args.range_stop_symbol("x", JDim).name: a.shape[1], + gtx_dace_args.field_stride_symbol("x", IDim).name: a.strides[0] // a.itemsize, + gtx_dace_args.field_stride_symbol("x", JDim).name: a.strides[1] // a.itemsize, + gtx_dace_args.range_start_symbol("y", IDim).name: 0, + gtx_dace_args.range_stop_symbol("y", IDim).name: b.shape[0], + gtx_dace_args.range_start_symbol("y", JDim).name: 0, + gtx_dace_args.range_stop_symbol("y", JDim).name: b.shape[1], + gtx_dace_args.field_stride_symbol("y", IDim).name: b.strides[0] // b.itemsize, + gtx_dace_args.field_stride_symbol("y", JDim).name: b.strides[1] // b.itemsize, + gtx_dace_args.range_start_symbol("w", IDim).name: 0, + gtx_dace_args.range_stop_symbol("w", IDim).name: c.shape[0], + gtx_dace_args.range_start_symbol("w", JDim).name: 0, + gtx_dace_args.range_stop_symbol("w", JDim).name: c.shape[1], + gtx_dace_args.field_stride_symbol("w", IDim).name: c.strides[0] // c.itemsize, + gtx_dace_args.field_stride_symbol("w", JDim).name: c.strides[1] // c.itemsize, + gtx_dace_args.range_start_symbol("z", IDim).name: 0, + gtx_dace_args.range_stop_symbol("z", IDim).name: d.shape[0], + gtx_dace_args.range_start_symbol("z", JDim).name: 0, + gtx_dace_args.range_stop_symbol("z", JDim).name: d.shape[1], + gtx_dace_args.field_stride_symbol("z", IDim).name: d.strides[0] // d.itemsize, + gtx_dace_args.field_stride_symbol("z", JDim).name: d.strides[1] // d.itemsize, } sdfg = build_dace_sdfg(testee, CARTESIAN_OFFSETS) @@ -2340,18 +2406,18 @@ def test_gtir_scan(id, use_symbolic_column_size): ref = np.add.accumulate(a, axis=1) + VAL symbols = FSYMBOLS | { - "__x_KDim_range_0": 0, - "__x_KDim_range_1": a.shape[1], - "__x_IDim_stride": a.strides[0] // a.itemsize, - "__x_KDim_stride": a.strides[1] // a.itemsize, - "__y_KDim_range_0": 0, - "__y_KDim_range_1": b.shape[1], - "__y_IDim_stride": b.strides[0] // b.itemsize, - "__y_KDim_stride": b.strides[1] // b.itemsize, - "__z_KDim_range_0": 0, - "__z_KDim_range_1": z.shape[1], - "__z_IDim_stride": z.strides[0] // z.itemsize, - "__z_KDim_stride": z.strides[1] // z.itemsize, + gtx_dace_args.range_start_symbol("x", KDim).name: 0, + gtx_dace_args.range_stop_symbol("x", KDim).name: a.shape[1], + gtx_dace_args.field_stride_symbol("x", IDim).name: a.strides[0] // a.itemsize, + gtx_dace_args.field_stride_symbol("x", KDim).name: a.strides[1] // a.itemsize, + gtx_dace_args.range_start_symbol("y", KDim).name: 0, + gtx_dace_args.range_stop_symbol("y", KDim).name: b.shape[1], + gtx_dace_args.field_stride_symbol("y", IDim).name: b.strides[0] // b.itemsize, + gtx_dace_args.field_stride_symbol("y", KDim).name: b.strides[1] // b.itemsize, + gtx_dace_args.range_start_symbol("z", KDim).name: 0, + gtx_dace_args.range_stop_symbol("z", KDim).name: z.shape[1], + gtx_dace_args.field_stride_symbol("z", IDim).name: z.strides[0] // z.itemsize, + gtx_dace_args.field_stride_symbol("z", KDim).name: z.strides[1] // z.itemsize, } sdfg(a, b, z, **symbols) @@ -2404,18 +2470,18 @@ def test_gtir_scan_single_level_output(): ref = np.add.accumulate(a, axis=1) symbols = FSYMBOLS | { - "__x_KDim_range_0": 0, - "__x_KDim_range_1": a.shape[1], - "__x_IDim_stride": a.strides[0] // a.itemsize, - "__x_KDim_stride": a.strides[1] // a.itemsize, - "__y_KDim_range_0": 0, - "__y_KDim_range_1": b.shape[1], - "__y_IDim_stride": b.strides[0] // b.itemsize, - "__y_KDim_stride": b.strides[1] // b.itemsize, - "__z_KDim_range_0": 0, - "__z_KDim_range_1": c.shape[1], - "__z_IDim_stride": c.strides[0] // c.itemsize, - "__z_KDim_stride": c.strides[1] // c.itemsize, + gtx_dace_args.range_start_symbol("x", KDim).name: 0, + gtx_dace_args.range_stop_symbol("x", KDim).name: a.shape[1], + gtx_dace_args.field_stride_symbol("x", IDim).name: a.strides[0] // a.itemsize, + gtx_dace_args.field_stride_symbol("x", KDim).name: a.strides[1] // a.itemsize, + gtx_dace_args.range_start_symbol("y", KDim).name: 0, + gtx_dace_args.range_stop_symbol("y", KDim).name: b.shape[1], + gtx_dace_args.field_stride_symbol("y", IDim).name: b.strides[0] // b.itemsize, + gtx_dace_args.field_stride_symbol("y", KDim).name: b.strides[1] // b.itemsize, + gtx_dace_args.range_start_symbol("z", KDim).name: 0, + gtx_dace_args.range_stop_symbol("z", KDim).name: c.shape[1], + gtx_dace_args.field_stride_symbol("z", IDim).name: c.strides[0] // c.itemsize, + gtx_dace_args.field_stride_symbol("z", KDim).name: c.strides[1] // c.itemsize, } sdfg(a, b, c, **symbols) diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index a6aad1cb56..5273c8262b 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -849,19 +849,23 @@ class TestCodegenName: ("a_b.c", "a_ub_dc"), ("my__mod.X", "my_u_umod_dX"), ("_CONST_DIM", "_uCONST_uDIM"), + # a parametrized tag: the brackets must not survive into the identifier + ("gt4py.next.common.Staggered[pkg.K]", "gt4py_dnext_dcommon_dStaggered_lpkg_dK_r"), ], ) def test_known_values(self, tag, expected): assert common.codegen_name(tag) == expected assert common.from_codegen_name(expected) == tag - @pytest.mark.parametrize("tag", ["_u", "_d", "a_ud.b", "_ud_du", "..", "__"]) + @pytest.mark.parametrize( + "tag", ["_u", "_d", "_l", "_r", "a_ud.b", "_ud_du", "..", "__", "[]", "a[b.c]", "_l[_r]"] + ) def test_roundtrip_adversarial(self, tag): """Tags that look like the escape sequences themselves must still round-trip.""" assert common.from_codegen_name(common.codegen_name(tag)) == tag def test_output_is_a_valid_identifier(self): - for tag in ["mod.V2E.Local", "a.b_c", "_CONST_DIM", "pkg.sub.Dim"]: + for tag in ["mod.V2E.Local", "a.b_c", "_CONST_DIM", "pkg.sub.Dim", "mod.Staggered[mod.K]"]: assert re.fullmatch(r"[A-Za-z_]\w*", common.codegen_name(tag)), tag def test_injective_and_reversible_exhaustively(self): @@ -872,9 +876,10 @@ def test_injective_and_reversible_exhaustively(self): The naive scheme -- `_` -> `__` then `.` -> `_` -- fails this with 686 collisions, because a dot becomes a single underscore and `'..'` collides with an escaped `'_'`. """ - alphabet = "a._ud" + # every character a tag can contain that needs escaping, plus the escape letters + alphabet = "a._[]udlr" seen: dict[str, str] = {} - for length in range(1, 6): + for length in range(1, 5): for tag in map("".join, itertools.product(alphabet, repeat=length)): mangled = common.codegen_name(tag) assert mangled not in seen, ( @@ -882,4 +887,4 @@ def test_injective_and_reversible_exhaustively(self): ) seen[mangled] = tag assert common.from_codegen_name(mangled) == tag - assert len(seen) == sum(len(alphabet) ** n for n in range(1, 6)) + assert len(seen) == sum(len(alphabet) ** n for n in range(1, 5)) From 5e9e5422f1412dcf8008b177818e74e4ec9fc7e1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 13:49:51 +0200 Subject: [PATCH 07/20] wip[next]: fingerprint staggered dims; IR text accepts tags; last gtfn names * Strict fingerprinter: `Staggered[K]` has no importable qualified name, so the by-reference `type` deconstruction rejected it. Registered on `StaggeredMeta`, it is fingerprinted by its base dimension (importable), the same reduction its `copyreg` hook uses; the bare `Staggered` stays by reference. * IR text format: the pretty parser read axis and offset literals as `CNAME`, so a qualified tag (dots, and brackets for `Staggered[pkg.K]`) did not round-trip. A `TAG` terminal describes a qualified name; unambiguous, since a tag starts with a letter and the literal's suffix terminates it. * gtfn: the scan column axis and sparse-argument tuple-like dimensions still named `generated::_t` with the raw tag. * The sparse-argument path read `.tag` off a legacy `FieldOffset`, which only has `.value` -- a leftover of the #2844 `.value -> .tag` replay. --- src/gt4py/next/fingerprinting.py | 11 ++++++++- src/gt4py/next/iterator/pretty_parser.py | 9 +++++-- .../codegens/gtfn/gtfn_module.py | 6 +++-- .../codegens/gtfn/itir_to_gtfn_ir.py | 24 ++++++++++++++----- 4 files changed, 39 insertions(+), 11 deletions(-) diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index f0739a102b..b762da1ec9 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -55,7 +55,7 @@ import xxhash from gt4py.eve import concepts, datamodels, utils as eve_utils -from gt4py.next import utils as next_utils +from gt4py.next import common, utils as next_utils _T = TypeVar("_T") @@ -222,6 +222,15 @@ def _instance_state(obj: Any) -> tuple[Any, ...]: *((member.name, member.value) for member in obj), state=b"enum_class\0" + eve_utils.get_fully_qualified_name(obj).encode(), ), + # A parametrized dimension such as `Staggered[K]` has no importable qualified name (its + # `__qualname__` contains brackets), so the by-reference `type` deconstruction rejects it. + # It is fully determined by its base dimension, which *is* importable -- the same reduction + # its `copyreg` registration uses. The bare `Staggered` base is an ordinary class. See ADR 0028. + common.StaggeredMeta: lambda obj: ( + Deconstruction.from_pieces(obj.base, state=b"staggered_dimension") + if "base" 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/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index b58f66436f..07033ff1db 100644 --- a/src/gt4py/next/iterator/pretty_parser.py +++ b/src/gt4py/next/iterator/pretty_parser.py @@ -35,8 +35,13 @@ TYPE_LITERAL: CNAME INT_LITERAL: SIGNED_INT FLOAT_LITERAL: SIGNED_FLOAT - OFFSET_LITERAL: ( INT_LITERAL | CNAME ) "ₒ" - AXIS_LITERAL: CNAME ("ᵥ" | "ₕ") + // A dimension or offset tag is a qualified Python name (ADR 0028): dotted, and -- for a + // parametrized dimension such as `Staggered[pkg.K]` -- with one bracketed dotted name. + // Unambiguous here: a tag starts with a letter (a float does not), and the literal's + // suffix terminates it. + TAG: /[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*(?:\[[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*\])?/ + OFFSET_LITERAL: ( INT_LITERAL | TAG ) "ₒ" + AXIS_LITERAL: TAG ("ᵥ" | "ₕ") INFINITY_LITERAL: "∞" | "-∞" _literal: INT_LITERAL | FLOAT_LITERAL | OFFSET_LITERAL | AXIS_LITERAL | INFINITY_LITERAL ID_NAME: CNAME 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 965b376c94..68591309c6 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -89,11 +89,13 @@ def _process_regular_arguments( or dim.kind is common.DimensionKind.LOCAL ): # translate sparse dimensions to tuple dtype - dim_name = dim.tag + # 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) size = connectivity.max_neighbors - arg = f"gridtools::sid::dimension_to_tuple_like({arg})" + arg = f"gridtools::sid::dimension_to_tuple_like({arg})" arg_exprs.append(arg) return parameters, arg_exprs 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 7910632191..74684d420b 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 @@ -107,7 +107,9 @@ def _collect_dimensions_from_domain( for nr in domain.args: assert isinstance(nr, itir.FunCall) dim_name = _name_from_named_range(nr) - offset_definitions[dim_name] = TagDefinition(name=Sym(id=common.codegen_name(dim_name))) + offset_definitions[dim_name] = TagDefinition( + name=Sym(id=common.codegen_name(dim_name)) + ) elif domain.fun == itir.SymRef(id="unstructured_domain"): if len(domain.args) > 2: raise ValueError("Unstructured_domain must not have more than 2 arguments.") @@ -148,7 +150,9 @@ def _collect_dimensions_from_params( for type_ in type_info.primitive_constituents(param.type): if isinstance(type_, ts.FieldType): for dim in type_.dims: - offset_definitions[dim.tag] = TagDefinition(name=Sym(id=common.codegen_name(dim.tag))) + offset_definitions[dim.tag] = TagDefinition( + name=Sym(id=common.codegen_name(dim.tag)) + ) return offset_definitions @@ -170,7 +174,9 @@ def _collect_offset_definitions( ] for dim in dims: if grid_type == common.GridType.CARTESIAN: - offset_definitions[dim.tag] = TagDefinition(name=Sym(id=common.codegen_name(dim.tag))) + offset_definitions[dim.tag] = TagDefinition( + name=Sym(id=common.codegen_name(dim.tag)) + ) else: assert grid_type == common.GridType.UNSTRUCTURED if dim.kind != common.DimensionKind.VERTICAL: @@ -184,7 +190,9 @@ def _collect_offset_definitions( for offset_name, connectivity_type in offset_provider_type.items(): if isinstance(connectivity_type, common.NeighborConnectivityType): assert grid_type == common.GridType.UNSTRUCTURED - offset_definitions[offset_name] = TagDefinition(name=Sym(id=common.codegen_name(offset_name))) + 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)) @@ -215,7 +223,10 @@ def _add_staggered_aliases( if tag_def.alias is None and (base_name := common.staggered_base_tag(name)) is not None: # ensure the base tag exists (as alias target and loop dimension) in this position result.setdefault(base_name, TagDefinition(name=Sym(id=common.codegen_name(base_name)))) - aliases[name] = TagDefinition(name=Sym(id=common.codegen_name(name)), alias=SymRef(id=common.codegen_name(base_name))) + aliases[name] = TagDefinition( + name=Sym(id=common.codegen_name(name)), + alias=SymRef(id=common.codegen_name(base_name)), + ) else: result[name] = tag_def return {**result, **aliases} @@ -681,7 +692,8 @@ def convert_el_to_sid(el_expr: Expr, el_type: ts.ScalarType | ts.FieldType) -> E backend=backend, scans=[scan], args=[self._visit_output_argument(node.target), *lowered_inputs], - axis=SymRef(id=column_axis.tag), + # the column axis names its `generated::_t` tag type: mangle as declared + axis=SymRef(id=common.codegen_name(column_axis.tag)), ) assert projector is None # only scans have projectors return StencilExecution( From 2c68d48299f3a49b3e404044bd907ab06d87379c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 14:29:20 +0200 Subject: [PATCH 08/20] wip[next]: notebooks under nominal identity; docs, examples and typing Everything now passes: unit 1764 + 451 (DaCe), integration and regression 2745 + DaCe, and every notebook `test_examples` runs. * Notebooks. Compilation runs in `spawn` workers by default, and a class declared in an interactive `__main__` (a notebook, the REPL, `python -c`) pickles in the parent but cannot be unpickled in a worker, which has no such `__main__` to import. Reproduced with CI's own mechanism (`nbmake`): "Can't get attribute 'IDim' on ". The project's own Quickstart and workshop do exactly this, so a documented limitation was not enough. The process runner now falls back to the calling thread, with a warning, when a job references a class from an interactive `__main__` -- the same fallback it already takes for an unpicklable executor. Scripts keep parallel compilation: spawn re-imports a script's `__main__`. ADR 0028 updated. * `__gt_dims__` is the interop protocol with `gt4py.cartesian`, which names axes by bare name ("I", "J", "K"), so it returns `__qualname__`, not the qualified tag -- otherwise cartesian transposes the array wrongly. No test covered it; one does now, and fails with the fix reverted. * Docs, workshop notebooks and `examples/` migrated; notebook code cells only, stored outputs untouched. Both the Quickstart and `slides_2` declared the cell dimension twice, harmless while declarations compared equal; the second one is dropped, since the Quickstart later combined fields built on each. * mypy clean on `src/`. `_is_field_axis` called `isinstance` on a PEP 695 alias, which raises at runtime; only an `assert` reaches it. `TYPE_BUILTINS` keeps the `common.Dimension` alias: its `__name__` is "Dimension", the DSL builtin's name. (#2844 swapped in `DimensionIndex`, which would rename the builtin.) * DaCe orchestration test and two diagnostics that joined tags into messages. --- .../next/0028-Dimensions_As_Nominal_Types.md | 22 +++++-- docs/user/next/QuickstartGuide.md | 31 +++++----- .../exercises/1_simple_addition.ipynb | 8 ++- .../1_simple_addition_solution.ipynb | 8 ++- docs/user/next/workshop/exercises/helpers.py | 57 +++++++++++++------ docs/user/next/workshop/slides/slides_1.ipynb | 7 ++- docs/user/next/workshop/slides/slides_2.ipynb | 19 ++++--- docs/user/next/workshop/slides/slides_3.ipynb | 6 +- docs/user/next/workshop/slides/slides_4.ipynb | 2 +- examples/lap_cartesian_vs_next.ipynb | 12 +++- src/gt4py/next/common.py | 18 +++--- src/gt4py/next/embedded/nd_array_field.py | 2 +- src/gt4py/next/ffront/fbuiltins.py | 2 +- .../ffront/foast_passes/type_deduction.py | 2 +- src/gt4py/next/iterator/embedded.py | 8 ++- src/gt4py/next/otf/runners.py | 54 ++++++++++++++++-- .../next/type_system/type_translation.py | 4 +- .../dace_tests/test_orchestration.py | 18 +++--- .../unit_tests/otf_tests/test_runners.py | 35 ++++++++++++ tests/next_tests/unit_tests/test_common.py | 12 ++++ 20 files changed, 243 insertions(+), 84 deletions(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 299f4a76b5..224f8ec289 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -93,11 +93,19 @@ disappears. a class deleted after creation passes), so the authoritative check remains pickle's own `save_global`. - **Known limitation**: interactive `__main__` — the REPL, notebooks, - `python -c` — cannot be resolved. The `spawn`-based compile workers - re-execute the main *script* as `__mp_main__`, so a dimension declared in a - file's `__main__` does resolve, provided the script has the - `if __name__ == "__main__":` guard the worker pool already requires. + **Interactive `__main__`** -- a notebook, the REPL, `python -c` -- has no + `__file__` for a worker to re-import, so a class declared there pickles in the + parent (which has it) and then fails to unpickle in a spawn worker. This is not + left as a documented limitation: the repository's own Quickstart and workshop + notebooks declare dimensions interactively and are run in CI. Instead the + process runner detects a job that references a class from an interactive + `__main__` and compiles it in the calling thread, with a warning -- the same + fallback it already takes for an executor that cannot be pickled. A *script's* + `__main__` is re-imported by spawn workers (as `__mp_main__`), so scripts keep + parallel compilation, provided they have the `if __name__ == "__main__":` guard + the worker pool already requires. `(tag, kind)` value identity with a registry + avoided this by pickling dimensions by value; that is the one case where it + was strictly more convenient. 3. **No registry and no blanket `copyreg`.** Module-level classes pickle by reference, which is correct. A *narrow* `copyreg` registration is still @@ -120,7 +128,9 @@ disappears. def from_codegen_name(name: str) -> Tag: - return re.sub(r"_([udlr])", lambda m: {"u": "_", "d": ".", "l": "[", "r": "]"}[m.group(1)], name) + return re.sub( + r"_([udlr])", lambda m: {"u": "_", "d": ".", "l": "[", "r": "]"}[m.group(1)], name + ) ``` A *prefix* escape, not `_ -> __` followed by `. -> _`: the latter is **not diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 829acfeffa..1e1cbfc280 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -51,11 +51,11 @@ from gt4py.next import float64, neighbor_sum, where, Dims #### Fields -Fields store data as a multi-dimensional array, and are defined over a set of named dimensions. The code snippet below defines two named dimensions, _Cell_ and _K_, and creates the fields `a` and `b` over their cartesian product using the `gtx.as_field` helper function. The fields contain the values 2 for `a` and 3 for `b` for all entries. +Fields store data as a multi-dimensional array, and are defined over a set of named dimensions. The code snippet below defines two dimensions, `CellDim` and `KDim` -- a dimension is a class -- and creates the fields `a` and `b` over their cartesian product using the `gtx.as_field` helper function. The fields contain the values 2 for `a` and 3 for `b` for all entries. ```{code-cell} ipython3 -CellDim = gtx.Dimension("Cell") -KDim = gtx.Dimension("K") +class CellDim(gtx.DimensionIndex): ... +class KDim(gtx.DimensionIndex): ... num_cells = 5 num_layers = 6 @@ -70,8 +70,8 @@ b = gtx.as_field([CellDim, KDim], np.full(shape=grid_shape, fill_value=b_value, Additional numpy-equivalent constructors are available, namely `ones`, `zeros`, `empty`, `full`. These require domain, dtype, and allocator (e.g. a backend) specifications. ```{code-cell} ipython3 -I = gtx.Dimension("I") -J = gtx.Dimension("J") +class I(gtx.DimensionIndex): ... +class J(gtx.DimensionIndex): ... array_of_ones_numpy = np.ones((grid_shape[0], grid_shape[1])) field_of_ones = gtx.ones( @@ -165,11 +165,10 @@ The examples related to unstructured meshes use the mesh below. The edges (in bl +++ -The fields in the subsequent code snippets are 1-dimensional, either over the cells or over the edges. The corresponding named dimensions are thus the following: +The fields in the subsequent code snippets are 1-dimensional, either over the cells or over the edges. They use the `CellDim` declared above, plus a new dimension for the edges. A dimension is identified by its class, so declaring `CellDim` again here would create a *different* dimension from the one the fields above are defined on: ```{code-cell} ipython3 -CellDim = gtx.Dimension("Cell") -EdgeDim = gtx.Dimension("Edge") +class EdgeDim(gtx.DimensionIndex): ... ``` You can express connectivity between elements (i.e. cells or edges) of the mesh using connectivity (a.k.a. adjacency or neighborhood) tables. The table below, `edge_to_cell_table`, has one row for every edge where it lists the indices of cells adjacent to that edge. For example, this table says that edge #6 connects to cells #0 and #5. Similarly, `cell_to_edge_table` lists the edges that are neighbors to a particular cell. @@ -228,11 +227,11 @@ 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 -E2CDim = gtx.Dimension("E2C", kind=gtx.DimensionKind.LOCAL) -E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim,E2CDim)) +class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim)) ``` -Note that the field offset does not contain the actual connectivity table, that's provided through an _offset provider_: +The field offset is named by its local dimension's `tag`, and the offset provider below is keyed by the same `tag`, so all three refer to one connectivity. Note that the field offset does not contain the actual connectivity table, that's provided through an _offset provider_: ```{code-cell} ipython3 E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2CDim], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) @@ -251,7 +250,7 @@ def nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64]) -> gtx. def run_nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): nearest_cell_to_edge(cell_values, out=out) -run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={"E2C": E2C_offset_provider}) +run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) print("0th adjacent cell's value: {}".format(edge_values.asnumpy())) ``` @@ -278,7 +277,7 @@ def sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64]) -> gtx.Field[D def run_sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): sum_adjacent_cells(cells, out=out) -run_sum_adjacent_cells(cell_values, edge_values, offset_provider={"E2C": E2C_offset_provider}) +run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) print("sum of adjacent cells: {}".format(edge_values.asnumpy())) ``` @@ -380,8 +379,8 @@ 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 -C2EDim = gtx.Dimension("C2E", kind=gtx.DimensionKind.LOCAL) -C2E = gtx.FieldOffset("C2E", source=EdgeDim, target=(CellDim, C2EDim)) +class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +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) ``` @@ -442,7 +441,7 @@ result_pseudo_lap = gtx.as_field([CellDim], np.zeros(shape=(6,))) run_pseudo_laplacian(cell_values, edge_weight_field, result_pseudo_lap, - offset_provider={"E2C": E2C_offset_provider, "C2E": C2E_offset_provider}) + offset_provider={E2CDim.tag: E2C_offset_provider, C2EDim.tag: C2E_offset_provider}) print("pseudo-laplacian: {}".format(result_pseudo_lap.asnumpy())) ``` diff --git a/docs/user/next/workshop/exercises/1_simple_addition.ipynb b/docs/user/next/workshop/exercises/1_simple_addition.ipynb index 7f42f2b9d8..64cd1c55c0 100644 --- a/docs/user/next/workshop/exercises/1_simple_addition.ipynb +++ b/docs/user/next/workshop/exercises/1_simple_addition.ipynb @@ -36,8 +36,12 @@ "metadata": {}, "outputs": [], "source": [ - "I = gtx.Dimension(\"I\")\n", - "J = gtx.Dimension(\"J\")\n", + "class I(gtx.DimensionIndex): ...\n", + "\n", + "\n", + "class J(gtx.DimensionIndex): ...\n", + "\n", + "\n", "size = 10" ] }, diff --git a/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb b/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb index dd1a30cc6b..40409ad8bc 100644 --- a/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb +++ b/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb @@ -51,8 +51,12 @@ "metadata": {}, "outputs": [], "source": [ - "I = gtx.Dimension(\"I\")\n", - "J = gtx.Dimension(\"J\")\n", + "class I(gtx.DimensionIndex): ...\n", + "\n", + "\n", + "class J(gtx.DimensionIndex): ...\n", + "\n", + "\n", "size = 10" ] }, diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index 8fd90daf20..c398524538 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -11,7 +11,7 @@ import gt4py.next as gtx from gt4py.next.iterator.embedded import MutableLocatedField from gt4py.next import neighbor_sum, where, Dims -from gt4py.next import Dimension, DimensionKind, FieldOffset +from gt4py.next import Dimension, DimensionIndex, DimensionKind, FieldOffset from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( run_gtfn as gtfn_cpu, @@ -377,18 +377,43 @@ def ripple_field(domain: gtx.Domain, *, allocator=None) -> MutableLocatedField: ) -C = Dimension("C") -V = Dimension("V") -E = Dimension("E") -K = Dimension("K", kind=gtx.DimensionKind.VERTICAL) - -C2EDim = Dimension("C2E", kind=DimensionKind.LOCAL) -C2E = FieldOffset("C2E", source=E, target=(C, C2EDim)) -V2EDim = Dimension("V2E", kind=DimensionKind.LOCAL) -V2E = FieldOffset("V2E", source=E, target=(V, V2EDim)) -E2VDim = Dimension("E2V", kind=DimensionKind.LOCAL) -E2V = FieldOffset("E2V", source=V, target=(E, E2VDim)) -E2CDim = Dimension("E2C", kind=DimensionKind.LOCAL) -E2C = FieldOffset("E2C", source=C, target=(E, E2CDim)) -E2C2VDim = Dimension("E2C2V", kind=DimensionKind.LOCAL) -E2C2V = FieldOffset("E2C2V", source=V, target=(E, E2C2VDim)) +class C(DimensionIndex): ... + + +class V(DimensionIndex): ... + + +class E(DimensionIndex): ... + + +class K(DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +C2E = FieldOffset(C2EDim.tag, source=E, target=(C, C2EDim)) + + +class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +V2E = FieldOffset(V2EDim.tag, source=E, target=(V, V2EDim)) + + +class E2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +E2V = FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) + + +class E2CDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +E2C = FieldOffset(E2CDim.tag, source=C, target=(E, E2CDim)) + + +class E2C2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +E2C2V = FieldOffset(E2C2VDim.tag, source=V, target=(E, E2C2VDim)) diff --git a/docs/user/next/workshop/slides/slides_1.ipynb b/docs/user/next/workshop/slides/slides_1.ipynb index f1e0229d4f..724162dc8f 100644 --- a/docs/user/next/workshop/slides/slides_1.ipynb +++ b/docs/user/next/workshop/slides/slides_1.ipynb @@ -207,8 +207,11 @@ } ], "source": [ - "Cell = gtx.Dimension(\"Cell\")\n", - "K = gtx.Dimension(\"K\", kind=gtx.DimensionKind.VERTICAL)\n", + "class Cell(gtx.DimensionIndex): ...\n", + "\n", + "\n", + "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ...\n", + "\n", "\n", "domain = gtx.domain({Cell: 5, K: 6})\n", "\n", diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index 6686bd7e14..db8f370abc 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -54,8 +54,10 @@ "metadata": {}, "outputs": [], "source": [ - "Cell = gtx.Dimension(\"Cell\")\n", - "K = gtx.Dimension(\"K\", kind=gtx.DimensionKind.VERTICAL)" + "class Cell(gtx.DimensionIndex): ...\n", + "\n", + "\n", + "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ..." ] }, { @@ -152,8 +154,7 @@ "metadata": {}, "outputs": [], "source": [ - "Cell = gtx.Dimension(\"Cell\")\n", - "Edge = gtx.Dimension(\"Edge\")" + "class Edge(gtx.DimensionIndex): ..." ] }, { @@ -272,8 +273,10 @@ "metadata": {}, "outputs": [], "source": [ - "E2CDim = gtx.Dimension(\"E2C\", kind=gtx.DimensionKind.LOCAL)\n", - "E2C = gtx.FieldOffset(\"E2C\", source=Cell, target=(Edge, E2CDim))" + "class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...\n", + "\n", + "\n", + "E2C = gtx.FieldOffset(E2CDim.tag, source=Cell, target=(Edge, E2CDim))" ] }, { @@ -317,7 +320,7 @@ " nearest_cell_to_edge(cell_field, out=edge_field)\n", "\n", "\n", - "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={\"E2C\": E2C_offset_provider})\n", + "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", "\n", "print(\"0th adjacent cell's value: {}\".format(edge_field.asnumpy()))" ] @@ -392,7 +395,7 @@ " sum_adjacent_cells(cell_field, out=edge_field)\n", "\n", "\n", - "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={\"E2C\": E2C_offset_provider})\n", + "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", "\n", "print(\"sum of adjacent cells: {}\".format(edge_field.asnumpy()))" ] diff --git a/docs/user/next/workshop/slides/slides_3.ipynb b/docs/user/next/workshop/slides/slides_3.ipynb index bd85f5027b..1d699e996f 100644 --- a/docs/user/next/workshop/slides/slides_3.ipynb +++ b/docs/user/next/workshop/slides/slides_3.ipynb @@ -54,8 +54,10 @@ "metadata": {}, "outputs": [], "source": [ - "Cell = gtx.Dimension(\"Cell\")\n", - "K = gtx.Dimension(\"K\", kind=gtx.DimensionKind.VERTICAL)" + "class Cell(gtx.DimensionIndex): ...\n", + "\n", + "\n", + "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ..." ] }, { diff --git a/docs/user/next/workshop/slides/slides_4.ipynb b/docs/user/next/workshop/slides/slides_4.ipynb index a5b8aaf78f..12870c30de 100644 --- a/docs/user/next/workshop/slides/slides_4.ipynb +++ b/docs/user/next/workshop/slides/slides_4.ipynb @@ -54,7 +54,7 @@ "metadata": {}, "outputs": [], "source": [ - "K = gtx.Dimension(\"K\", kind=gtx.DimensionKind.VERTICAL)" + "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ..." ] }, { diff --git a/examples/lap_cartesian_vs_next.ipynb b/examples/lap_cartesian_vs_next.ipynb index 9a8dfc92b5..9413ee57a4 100644 --- a/examples/lap_cartesian_vs_next.ipynb +++ b/examples/lap_cartesian_vs_next.ipynb @@ -65,10 +65,16 @@ "# allocator = gtx.gtfn_cpu\n", "# allocator = gtx.gtfn_gpu\n", "\n", + "\n", "# Note: for gt4py.next, names don't matter, for gt4py.cartesian they have to be \"I\", \"J\", \"K\"\n", - "I = gtx.Dimension(\"I\")\n", - "J = gtx.Dimension(\"J\")\n", - "K = gtx.Dimension(\"K\", kind=gtx.DimensionKind.VERTICAL)\n", + "class I(gtx.DimensionIndex): ...\n", + "\n", + "\n", + "class J(gtx.DimensionIndex): ...\n", + "\n", + "\n", + "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ...\n", + "\n", "\n", "domain = gtx.domain({I: nx, J: ny, K: nz})\n", "\n", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 983b702b67..be94869e63 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -10,9 +10,9 @@ import abc import collections +import copyreg import dataclasses import enum -import copyreg import functools import importlib import math @@ -212,6 +212,8 @@ def __eq__( # type: ignore[misc] # NOTE: dimension-vs-dimension comparison is deliberately *not* handled here. A # dimension's identity is its type, so `type.__eq__` (identity) is the correct # answer; overriding it with `(tag, kind)` equality is what ADR 0028 rejects. + if isinstance(value, DimensionMeta): + return NotImplemented # both sides decline, so Python falls back to identity if isinstance(value, core_defs.INTEGRAL_TYPES): return Domain(dims=(cls,), ranges=(UnitRange(value, value + 1),)) return NotImplemented @@ -389,7 +391,7 @@ def resolve(tag: Tag) -> Dimension: ) from ex if not isinstance(obj, DimensionMeta): raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") - return obj + return cast(Dimension, obj) raise ValueError( f"Cannot resolve dimension tag '{tag}': no importable module prefix. A dimension" " referenced from the IR must be declared at module level in an importable module." @@ -974,7 +976,10 @@ def __gt_domain__(self) -> Domain: @property def __gt_dims__(self) -> tuple[str, ...]: - return tuple(d.tag for d in self.__gt_domain__.dims) + # NOTE: the unqualified name, not the `tag`. This is the interop protocol with + # `gt4py.cartesian`, which identifies axes by their bare names (`"I"`, `"J"`, `"K"`); a + # qualified tag would not match and the axes would be transposed wrongly (ADR 0028). + return tuple(d.__qualname__ for d in self.__gt_domain__.dims) @runtime_checkable @@ -1715,8 +1720,7 @@ def tag(cls) -> Tag: def __getitem__(cls, base: Dimension) -> Dimension: if "base" in cls.__dict__: raise TypeError( - f"'{cls.__qualname__}' is already staggered; a dimension cannot be staggered" - " twice." + f"'{cls.__qualname__}' is already staggered; a dimension cannot be staggered twice." ) if not isinstance(base, DimensionMeta): raise TypeError(f"'Staggered' expects a dimension, got '{base!r}'.") @@ -1747,7 +1751,7 @@ def __getitem__(cls, base: Dimension) -> Dimension: # inside `Field[Dims[Staggered[K]], ...]`. The runtime form below builds a real, interned # class so that `issubclass` and eve's `type[...]` validation work. Verified clean under # `mypy --strict` and pyright. - class Staggered[D: DimensionIndex](DimensionIndex): # noqa: D101 [undocumented-public-class] + class Staggered[D: DimensionIndex](DimensionIndex): base: ClassVar[Dimension] else: @@ -1811,7 +1815,7 @@ def _make_staggered(base: Dimension) -> Dimension: return Staggered[base] # type: ignore[valid-type] # runtime subscription, see StaggeredMeta -copyreg.pickle(StaggeredMeta, _reduce_staggered) # type: ignore[arg-type] # metaclass reducer +copyreg.pickle(StaggeredMeta, _reduce_staggered) def is_staggered(dim: Dimension) -> bool: diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index 916cccfb67..b8133cbafa 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -934,7 +934,7 @@ def _concat_where( def _as_offset(offset: fbuiltins.FieldOffset, offset_field: NdArrayField) -> common.Connectivity: if not fbuiltins.is_cartesian_offset(offset): - target_dims = ", ".join(d.tag for d in offset.target) + target_dims = ", ".join(d.__qualname__ for d in offset.target) # for the diagnostic raise ValueError( f"'as_offset' is only supported for Cartesian offsets " f"(single target dimension equal to source dimension); " diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 682af608ad..7474cd1406 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -96,7 +96,7 @@ PYTHON_TYPE_BUILTINS = [bool, int, float, tuple] PYTHON_TYPE_BUILTIN_NAMES = [t.__name__ for t in PYTHON_TYPE_BUILTINS] -TYPE_BUILTINS = [ +TYPE_BUILTINS: list[Any] = [ common.Field, common.Dimension, int8, diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index cc8fe1a52c..3b20fbfb86 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -989,7 +989,7 @@ def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: assert isinstance(arg_0, ts.OffsetType) assert isinstance(arg_1, ts.FieldType) if not fbuiltins.is_cartesian_offset(arg_0): - target_dims = ", ".join(d.tag for d in arg_0.target) + target_dims = ", ".join(d.__qualname__ for d in arg_0.target) # for the diagnostic raise errors.DSLError( node.location, f"'as_offset' is only supported for Cartesian offsets " diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 9b2179e243..766b6b2726 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -1016,7 +1016,9 @@ def __gt_origin__(self) -> tuple[int, ...]: def _is_field_axis(axis: Axis) -> TypeGuard[FieldAxis]: - return isinstance(axis, FieldAxis) + # `FieldAxis` is `common.Dimension`, a PEP 695 alias, which `isinstance` rejects; a + # dimension is a class, i.e. an instance of its metaclass. + return isinstance(axis, common.DimensionMeta) def _is_tuple_axis(axis: Axis) -> TypeGuard[TupleAxis]: @@ -1156,7 +1158,9 @@ def premap( def restrict(self, item: common.AnyIndexSpec) -> Self: if isinstance(item, Sequence) and all(isinstance(e, common.DimensionIndex) for e in item): assert len(item) == 1 - assert isinstance(item[0], common.DimensionIndex) # for mypy errors on multiple lines below + assert isinstance( + item[0], common.DimensionIndex + ) # for mypy errors on multiple lines below # an index is a `DimensionIndex` instance now, not a (dim, value) namedtuple d, r = item[0].dim, item[0].value assert d == self._dimension diff --git a/src/gt4py/next/otf/runners.py b/src/gt4py/next/otf/runners.py index f260a33061..75b54f19a3 100644 --- a/src/gt4py/next/otf/runners.py +++ b/src/gt4py/next/otf/runners.py @@ -13,6 +13,7 @@ import atexit import concurrent.futures import dataclasses +import io import multiprocessing import os import pathlib @@ -180,15 +181,48 @@ def _run_compilation_task_in_worker( return executor(compilable) +def _interactive_main_reference(obj: object) -> str | None: + """ + Return the name of a class from an interactive `__main__` that `obj` references, if any. + + A spawn worker re-imports a *script's* `__main__` (as `__mp_main__`), so classes declared + there resolve in the worker. An interactive `__main__` -- a notebook kernel, the REPL, + `python -c` -- has no `__file__` to re-import, so a class declared there pickles fine in the + parent (which has it) and then fails to unpickle in the worker. Dimensions are classes + identified by their qualified name (ADR 0028), which makes this the common case in + notebooks. + + Only scans when `__main__` is interactive, so ordinary scripts pay nothing. + """ + main = sys.modules.get("__main__") + if main is None or getattr(main, "__file__", None): + return None + found: list[str] = [] + + class _Scanner(pickle.Pickler): + # `reducer_override` is consulted for classes too, before the by-reference default + def reducer_override(self, o: object) -> Any: + if not found and isinstance(o, type) and o.__module__ == "__main__": + found.append(o.__qualname__) + return NotImplemented + + try: + _Scanner(io.BytesIO()).dump(obj) + except Exception: # an unpicklable object is reported by the executor check instead + return None + return found[0] if found else None + + class ProcessRunner: """Compiles in a ``ProcessPoolExecutor`` (``spawn``). The worker runs the task's executor (post-lowering compile) and returns the picklable ``CompilationArtifact``. - Tasks that cannot be offloaded — a known ``no_offload_reason`` or an executor - that stdlib ``pickle`` cannot serialize — are compiled in the calling thread - instead (with a warning), so they behave as under ``SerialRunner``. + Tasks that cannot be offloaded — a known ``no_offload_reason``, an executor + that stdlib ``pickle`` cannot serialize, or a reference to a class declared in an + interactive ``__main__`` — are compiled in the calling thread instead (with a + warning), so they behave as under ``SerialRunner``. """ def __init__(self, max_workers: int, shared_session_cache_dir: str) -> None: @@ -218,7 +252,16 @@ def submit( executor_blob = pickle.dumps(task.executor) except Exception as error: # pickling arbitrary object graphs raises arbitrary errors reason = f"its executor is not picklable ({error!s})" - if executor_blob is None: + compilable = task.construct_compilable(True) + if ( + reason is None + and (name := _interactive_main_reference((task.executor, compilable))) is not None + ): + reason = ( + f"it references '{name}', declared in an interactive '__main__' (a notebook, the" + " REPL or 'python -c'), which a worker process cannot import" + ) + if reason is not None: warnings.warn( f"Compiling '{task.name}' in the calling thread instead of a worker process " f"because {reason}.", @@ -226,10 +269,11 @@ def submit( ) return _run_in_calling_thread(task) + assert executor_blob is not None # set whenever no reason to fall back was found return self._pool.submit( _run_compilation_task_in_worker, executor_blob=executor_blob, - compilable=task.construct_compilable(True), + compilable=compilable, config_overrides=_config_snapshot(), recursion_limit=sys.getrecursionlimit(), ) diff --git a/src/gt4py/next/type_system/type_translation.py b/src/gt4py/next/type_system/type_translation.py index aa7730b97d..9722b27f5b 100644 --- a/src/gt4py/next/type_system/type_translation.py +++ b/src/gt4py/next/type_system/type_translation.py @@ -19,7 +19,7 @@ import sys import types import typing -from typing import Any, ForwardRef, Optional, TypeAlias +from typing import Any, ForwardRef, Optional, TypeAlias, cast import numpy as np import numpy.typing as npt @@ -221,7 +221,7 @@ def from_type_hint( for d in dim_arg: if not isinstance(d, common.DimensionMeta): raise ValueError(f"Invalid field dimension definition '{d}'.") - dims.append(d) + dims.append(cast(common.Dimension, d)) else: raise ValueError(f"Invalid field dimensions '{dim_arg}'.") diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py index 7dc8943e81..47551e473c 100644 --- a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py +++ b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py @@ -11,9 +11,10 @@ import pytest import gt4py.next as gtx +from gt4py._core import definitions as core_defs from gt4py.next import common as gtx_common, custom_layout_allocators as gtx_allocators +from gt4py.next.program_processors.runners.dace import sdfg_args as gtx_dace_args -from gt4py._core import definitions as core_defs from next_tests.integration_tests import cases from next_tests.integration_tests.cases import cartesian_case, unstructured_case # noqa: F401 from next_tests.integration_tests.cases_utils import ( @@ -80,6 +81,7 @@ def testee(a: gtx.Field[gtx.Dims[Vertex], gtx.float64], b: gtx.Field[gtx.Dims[Ed @pytest.mark.uses_unstructured_shift def test_sdfgConvertible_connectivities(unstructured_case): # noqa: F811 + E2V_CONN = gtx_dace_args.connectivity_identifier(E2VDim.tag) if not unstructured_case.backend or "dace" not in unstructured_case.backend.name: pytest.skip("DaCe-related test: Test SDFGConvertible interface for GT4Py programs") @@ -148,9 +150,10 @@ def get_stride_from_numpy_to_dace(arg: core_defs.NDArrayObject, axis: int) -> in offset_provider, rows=3, cols=2, - gt_conn_E2V=e2v, - __gt_conn_E2V_source_stride=get_stride_from_numpy_to_dace(e2v.ndarray, 0), - __gt_conn_E2V_neighbor_stride=get_stride_from_numpy_to_dace(e2v.ndarray, 1), + # the connectivity argument is named after the mangled offset key (ADR 0028) + **{E2V_CONN: e2v}, + **{f"__{E2V_CONN}_source_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 0)}, + **{f"__{E2V_CONN}_neighbor_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 1)}, ) e2v_np = e2v.asnumpy() @@ -169,9 +172,10 @@ def get_stride_from_numpy_to_dace(arg: core_defs.NDArrayObject, axis: int) -> in offset_provider, rows=3, cols=2, - gt_conn_E2V=e2v, - __gt_conn_E2V_source_stride=get_stride_from_numpy_to_dace(e2v.ndarray, 0), - __gt_conn_E2V_neighbor_stride=get_stride_from_numpy_to_dace(e2v.ndarray, 1), + # the connectivity argument is named after the mangled offset key (ADR 0028) + **{E2V_CONN: e2v}, + **{f"__{E2V_CONN}_source_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 0)}, + **{f"__{E2V_CONN}_neighbor_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 1)}, ) e2v_np = e2v.asnumpy() 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 c3c6aa6c7b..2ea966ff7e 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_runners.py +++ b/tests/next_tests/unit_tests/otf_tests/test_runners.py @@ -346,3 +346,38 @@ def test_default_runner_is_serial_in_worker_process(): runners.reset_default_runner() assert isinstance(runners.get_default_runner(), runners.SerialRunner) runners.reset_default_runner() + + +class TestInteractiveMainReference: + """ + A class declared in an interactive `__main__` pickles in the parent but not in a worker. + + Dimensions are classes identified by their qualified name (ADR 0028), so a notebook that + declares one would otherwise break the default process-pool compilation. + """ + + @staticmethod + def _main_class(name: str) -> type: + cls = type(name, (), {}) + cls.__module__ = "__main__" + return cls + + def test_found_when_main_is_interactive(self, monkeypatch): + interactive_main = type(sys)("__main__") # a notebook / REPL `__main__`: no `__file__` + cls = self._main_class("NotebookDim") + interactive_main.NotebookDim = cls + monkeypatch.setitem(sys.modules, "__main__", interactive_main) + assert runners._interactive_main_reference({"nested": [cls]}) == "NotebookDim" + + def test_ignored_when_main_is_a_script(self, monkeypatch): + # a spawn worker re-imports a script's `__main__`, so its classes do resolve there + script_main = type(sys)("__main__") + script_main.__file__ = "/some/script.py" + cls = self._main_class("ScriptDim") + script_main.ScriptDim = cls + monkeypatch.setitem(sys.modules, "__main__", script_main) + assert runners._interactive_main_reference(cls) is None + + def test_ignored_for_classes_from_ordinary_modules(self, monkeypatch): + monkeypatch.setitem(sys.modules, "__main__", type(sys)("__main__")) + assert runners._interactive_main_reference([dataclasses.dataclass, int]) is None diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index 5273c8262b..dcdc6128d2 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -8,6 +8,8 @@ import itertools import operator + +import numpy as np from typing import Optional, Pattern import pytest @@ -888,3 +890,13 @@ def test_injective_and_reversible_exhaustively(self): seen[mangled] = tag assert common.from_codegen_name(mangled) == tag assert len(seen) == sum(len(alphabet) ** n for n in range(1, 5)) + + +def test_gt_dims_are_unqualified_names(): + """ + `__gt_dims__` is the interop protocol with `gt4py.cartesian`, which names axes by their bare + names (`"I"`, `"J"`, `"K"`). A dimension's `tag` is its qualified name (ADR 0028), which + cartesian would not recognize, and would then transpose the array wrongly. + """ + field = gtx.as_field([IDim, JDim], np.zeros((2, 3))) + assert field.__gt_dims__ == ("IDim", "JDim") From 263af8fd3e9dffed34f497fb99704ada9c5b6b2c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 15:56:56 +0200 Subject: [PATCH 09/20] wip[next]: migrate src doctests to dimension classes The `src/` doctests were the one suite not covered locally: `pytest tests/` does not collect them, and nox runs them last, after the main suite. The Python 3.13 and 3.14 nox sessions caught 40 failures there -- all `Dimension("I")` calls on what is now a non-callable alias. They pass now on 3.12, 3.13 and 3.14 (116). * 60 doctest declarations become `class X(DimensionIndex): ...`, and 8 inline `Dimension("I")` expressions get a declaration hoisted in front of them. * Expected outputs showing the old `Dimension(value='I', ...)` repr are updated to the actual output. Applied only where the difference is purely a dimension's spelling, checked mechanically; one update the check refused was a false positive and was verified by eye. * A class declared in a doctest runs in a *copy* of the module's globals, so it is not an attribute of the real module and `resolve()` cannot import it. Four doctests run an IR pass that resolves a dimension; they bind the class onto the module explicitly, with a comment -- the same consequence of nominal identity as interactive `__main__`, confined to doctests. * Those bindings are separate `>>>` statements: written as `import sys; ...` on one line, ruff's docstring formatter split them into a continuation that doctest compiles as a single statement. --- src/gt4py/next/common.py | 35 +++++++++-------- src/gt4py/next/constructors.py | 22 +++++------ src/gt4py/next/embedded/common.py | 8 ++-- src/gt4py/next/ffront/decorator.py | 2 +- .../ffront/foast_passes/type_deduction.py | 8 ++-- src/gt4py/next/ffront/foast_pretty_printer.py | 3 +- src/gt4py/next/ffront/foast_to_gtir.py | 7 +++- src/gt4py/next/ffront/foast_to_past.py | 2 +- src/gt4py/next/ffront/func_to_foast.py | 5 ++- src/gt4py/next/ffront/func_to_past.py | 2 +- src/gt4py/next/ffront/past_to_itir.py | 8 ++-- src/gt4py/next/ffront/type_info.py | 9 ++--- src/gt4py/next/field_utils.py | 6 ++- src/gt4py/next/iterator/embedded.py | 8 ++-- src/gt4py/next/iterator/ir_utils/ir_makers.py | 14 +++---- .../iterator/transforms/fuse_as_fieldop.py | 15 +++++-- .../iterator/transforms/inline_fundefs.py | 4 +- .../transforms/prune_empty_concat_where.py | 6 ++- .../iterator/transforms/remove_broadcast.py | 14 +++++-- ...replace_get_domain_range_with_constants.py | 7 ++-- .../iterator/type_system/type_synthesizer.py | 39 ++++++++++++------- src/gt4py/next/type_system/type_info.py | 24 +++++++----- 22 files changed, 147 insertions(+), 101 deletions(-) diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index be94869e63..b479a81326 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -759,16 +759,16 @@ def __and__(self, other: Domain) -> Domain: Intersect `Domain`s, missing `Dimension`s are considered infinite. Examples: - >>> I = Dimension("I") - >>> J = Dimension("J") + >>> class I(DimensionIndex): ... + >>> class J(DimensionIndex): ... >>> Domain(NamedRange(I, UnitRange(-1, 3))) & Domain(NamedRange(I, UnitRange(1, 6))) - Domain(dims=(Dimension(value='I', kind=),), ranges=(UnitRange(1, 3),)) + Domain(dims=(gt4py.next.common.I[horizontal],), ranges=(UnitRange(1, 3),)) >>> Domain(NamedRange(I, UnitRange(-1, 3)), NamedRange(J, UnitRange(2, 4))) & Domain( ... NamedRange(I, UnitRange(1, 6)) ... ) - Domain(dims=(Dimension(value='I', kind=), Dimension(value='J', kind=)), ranges=(UnitRange(1, 3), UnitRange(2, 4))) + Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(1, 3), UnitRange(2, 4))) """ broadcast_dims = tuple(promote_dims(self.dims, other.dims)) intersected_ranges = tuple( @@ -816,10 +816,11 @@ def slice_at(self) -> utils.IndexerCallable[slice, Domain]: Create a new domain by slicing the domain ranges at the provided relative slices. Examples: - >>> I, J = Dimension("I"), Dimension("J") + >>> class I(DimensionIndex): ... + >>> class J(DimensionIndex): ... >>> domain = Domain(NamedRange(I, UnitRange(0, 10)), NamedRange(J, UnitRange(5, 15))) >>> domain.slice_at[2:3, 2:5] - Domain(dims=(Dimension(value='I', kind=), Dimension(value='J', kind=)), ranges=(UnitRange(2, 3), UnitRange(7, 10))) + Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(2, 3), UnitRange(7, 10))) """ def _domain_slicer(*args: slice) -> Domain: @@ -907,20 +908,20 @@ def domain(domain_like: DomainLike) -> Domain: Construct `Domain` from `DomainLike` object. Examples: - >>> I = Dimension("I") - >>> J = Dimension("J") + >>> class I(DimensionIndex): ... + >>> class J(DimensionIndex): ... >>> domain(((I, (2, 4)), (J, (3, 5)))) - Domain(dims=(Dimension(value='I', kind=), Dimension(value='J', kind=)), ranges=(UnitRange(2, 4), UnitRange(3, 5))) + Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(2, 4), UnitRange(3, 5))) >>> domain({I: (2, 4), J: (3, 5)}) - Domain(dims=(Dimension(value='I', kind=), Dimension(value='J', kind=)), ranges=(UnitRange(2, 4), UnitRange(3, 5))) + Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(2, 4), UnitRange(3, 5))) >>> domain(((I, 2), (J, 4))) - Domain(dims=(Dimension(value='I', kind=), Dimension(value='J', kind=)), ranges=(UnitRange(0, 2), UnitRange(0, 4))) + Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(0, 2), UnitRange(0, 4))) >>> domain({I: 2, J: 4}) - Domain(dims=(Dimension(value='I', kind=), Dimension(value='J', kind=)), ranges=(UnitRange(0, 2), UnitRange(0, 4))) + Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(0, 2), UnitRange(0, 4))) """ if isinstance(domain_like, Domain): return domain_like @@ -1617,11 +1618,11 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: Examples: >>> from gt4py.next.common import Dimension - >>> I = Dimension("I", DimensionKind.HORIZONTAL) - >>> J = Dimension("J", DimensionKind.HORIZONTAL) - >>> K = Dimension("K", DimensionKind.VERTICAL) - >>> E2V = Dimension("E2V", kind=DimensionKind.LOCAL) - >>> E2C = Dimension("E2C", kind=DimensionKind.LOCAL) + >>> class I(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + >>> class J(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + >>> class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... + >>> class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... + >>> class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... >>> promote_dims([J, K], [I, K]) == [I, J, K] True >>> promote_dims([K, J], [I, K]) diff --git a/src/gt4py/next/constructors.py b/src/gt4py/next/constructors.py index cb20f79481..c07a085fbb 100644 --- a/src/gt4py/next/constructors.py +++ b/src/gt4py/next/constructors.py @@ -431,7 +431,7 @@ def empty( Initialize a field in one dimension with a backend and a range domain: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> a = gtx.empty({IDim: range(3, 10)}, allocator=gtx.itir_python) >>> a.shape (7,) @@ -440,7 +440,7 @@ def empty( >>> import numpy as np >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> a = gtx.empty({IDim: range(3, 10)}, allocator=np) >>> a.shape (7,) @@ -448,7 +448,7 @@ def empty( Initialize with a device and an integer domain. It works like a shape with named dimensions: >>> from gt4py._core import definitions as core_defs - >>> JDim = gtx.Dimension("J") + >>> class JDim(gtx.DimensionIndex): ... >>> b = gtx.empty( ... {IDim: 3, JDim: 3}, int, device=core_defs.Device(core_defs.DeviceType.CPU, 0) ... ) @@ -476,7 +476,7 @@ def zeros( Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> gtx.zeros({IDim: range(3, 10)}, allocator=gtx.itir_python).ndarray array([0., 0., 0., 0., 0., 0., 0.]) """ @@ -501,7 +501,7 @@ def ones( Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> gtx.ones({IDim: range(3, 10)}, allocator=gtx.itir_python).ndarray array([1., 1., 1., 1., 1., 1., 1.]) """ @@ -532,7 +532,7 @@ def full( Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> gtx.full({IDim: 3}, 5, allocator=gtx.itir_python).ndarray array([5, 5, 5]) """ @@ -577,7 +577,7 @@ def as_field( Examples: >>> import numpy as np >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> xdata = np.array([1, 2, 3]) Automatic domain from just dimensions: @@ -651,9 +651,9 @@ def as_connectivity( Examples: >>> import numpy as np >>> from gt4py import next as gtx - >>> Vertex = gtx.Dimension("Vertex") - >>> Edge = gtx.Dimension("Edge") - >>> V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) + >>> class Vertex(gtx.DimensionIndex): ... + >>> class Edge(gtx.DimensionIndex): ... + >>> class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... >>> data = np.array([[0, 1], [1, 2], [2, 0]]) >>> conn = gtx.as_connectivity([Vertex, V2EDim], Edge, data) >>> conn.ndarray @@ -661,7 +661,7 @@ def as_connectivity( [1, 2], [2, 0]]) >>> conn.domain - Domain(dims=(Dimension(value='Vertex', kind=), Dimension(value='V2E', kind=)), ranges=(UnitRange(0, 3), UnitRange(0, 2))) + Domain(dims=(gt4py.next.constructors.Vertex[horizontal], gt4py.next.constructors.V2EDim[local]), ranges=(UnitRange(0, 3), UnitRange(0, 2))) """ if skip_value is eve.NOTHING: skip_value = ( diff --git a/src/gt4py/next/embedded/common.py b/src/gt4py/next/embedded/common.py index 07402b1079..3374d30e36 100644 --- a/src/gt4py/next/embedded/common.py +++ b/src/gt4py/next/embedded/common.py @@ -103,11 +103,11 @@ def domain_intersection(*domains: common.Domain) -> common.Domain: Return the intersection of the given domains. Example: - >>> I = common.Dimension("I") + >>> class I(common.DimensionIndex): ... >>> domain_intersection( ... common.domain({I: (0, 5)}), common.domain({I: (1, 3)}) ... ) # doctest: +ELLIPSIS - Domain(dims=(Dimension(value='I', ...), ranges=(UnitRange(1, 3),)) + Domain(dims=(gt4py.next.embedded.common.I[horizontal],), ranges=(UnitRange(1, 3),)) """ return functools.reduce(operator.and_, domains, common.Domain(dims=tuple(), ranges=tuple())) @@ -120,8 +120,8 @@ def restrict_to_intersection( Return the with each other intersected domains, ignoring 'ignore_dims' dimensions for the intersection. Example: - >>> I = common.Dimension("I") - >>> J = common.Dimension("J") + >>> class I(common.DimensionIndex): ... + >>> class J(common.DimensionIndex): ... >>> res = restrict_to_intersection( ... common.domain({I: (0, 5), J: (1, 2)}), ... common.domain({I: (1, 3), J: (0, 3)}), diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index bf63aa692a..77d87409d4 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -844,7 +844,7 @@ def scan_operator( >>> import gt4py.next as gtx >>> from gt4py.next.iterator import embedded >>> embedded._column_range = 1 # implementation detail - >>> KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) + >>> class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... >>> inp = gtx.as_field([KDim], np.ones((10,))) >>> out = gtx.as_field([KDim], np.zeros((10,))) >>> @gtx.scan_operator(axis=KDim, forward=True, init=0.0) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 3b20fbfb86..b439b7d77e 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -41,9 +41,8 @@ def with_altered_scalar_kind( >>> print(with_altered_scalar_kind(scalar_t, ts.ScalarKind.BOOL)) bool - >>> field_t = ts.FieldType( - ... dims=[Dimension(value="I")], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64) - ... ) + >>> class I(common.DimensionIndex): ... + >>> field_t = ts.FieldType(dims=[I], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) >>> print(with_altered_scalar_kind(field_t, ts.ScalarKind.FLOAT32)) Field[[I], float32] """ @@ -174,7 +173,8 @@ class FieldOperatorTypeDeduction(traits.VisitorWithSymbolTableTrait, NodeTransla >>> from gt4py.next import Field >>> from gt4py.next.ffront.source_utils import SourceDefinition, get_closure_vars_from_function >>> from gt4py.next.ffront.func_to_foast import FieldOperatorParser - >>> IDim = Dimension("IDim") + >>> from gt4py.next.common import DimensionIndex + >>> class IDim(DimensionIndex): ... >>> def example(a: "Field[[IDim], float]", b: "Field[[IDim], float]"): ... return a + b diff --git a/src/gt4py/next/ffront/foast_pretty_printer.py b/src/gt4py/next/ffront/foast_pretty_printer.py index 8b2e369501..9e2745f0ac 100644 --- a/src/gt4py/next/ffront/foast_pretty_printer.py +++ b/src/gt4py/next/ffront/foast_pretty_printer.py @@ -238,7 +238,8 @@ def pretty_format(node: foast.LocatedNode) -> str: Pretty print (to string) an `foast.LocatedNode`. >>> from gt4py.next import Field, Dimension, field_operator, float64 - >>> IDim = Dimension("IDim") + >>> from gt4py.next.common import DimensionIndex + >>> class IDim(DimensionIndex): ... >>> @field_operator ... def field_op(a: Field[[IDim], float64]) -> Field[[IDim], float64]: ... return a + 1.0 diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 2b2553f814..4817072a26 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -80,7 +80,8 @@ class FieldOperatorLowering(eve.PreserveLocationVisitor, eve.NodeTranslator): >>> from gt4py.next.ffront.func_to_foast import FieldOperatorParser >>> from gt4py.next import Field, Dimension, float64 >>> - >>> IDim = Dimension("IDim") + >>> from gt4py.next.common import DimensionIndex + >>> class IDim(DimensionIndex): ... >>> def fieldop(inp: Field[[IDim], "float64"]): ... return inp >>> @@ -309,7 +310,9 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: # `field(Dim + idx)` (where `idx` is integer or half integer) case foast.BinOp( op=dialect_ast_enums.BinaryOperator.ADD | dialect_ast_enums.BinaryOperator.SUB, - left=foast.LocatedNode(type=ts.DimensionType(dim=common.DimensionMeta() as dim)), + left=foast.LocatedNode( + type=ts.DimensionType(dim=common.DimensionMeta() as dim) + ), right=foast.Constant(value=offset_index), ): if arg.op == dialect_ast_enums.BinaryOperator.SUB: diff --git a/src/gt4py/next/ffront/foast_to_past.py b/src/gt4py/next/ffront/foast_to_past.py index beabe53c71..64a10a85e8 100644 --- a/src/gt4py/next/ffront/foast_to_past.py +++ b/src/gt4py/next/ffront/foast_to_past.py @@ -63,7 +63,7 @@ class OperatorToProgram(workflow.Workflow[ConcreteFOASTOperatorDef, ConcretePAST Example: >>> from gt4py import next as gtx >>> from gt4py.next.otf import arguments, workflow - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> @gtx.field_operator ... def copy(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: diff --git a/src/gt4py/next/ffront/func_to_foast.py b/src/gt4py/next/ffront/func_to_foast.py index 14dceb25d1..8cc9618ab1 100644 --- a/src/gt4py/next/ffront/func_to_foast.py +++ b/src/gt4py/next/ffront/func_to_foast.py @@ -53,7 +53,7 @@ def func_to_foast(inp: DSLFieldOperatorDef) -> FOASTOperatorDef: Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> const = gtx.float32(2.0) >>> def dsl_operator(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: @@ -140,7 +140,8 @@ class FieldOperatorParser(DialectParser[foast.FunctionDefinition]): >>> from gt4py.next import Field, Dimension >>> float64 = float - >>> IDim = Dimension("IDim") + >>> from gt4py.next.common import DimensionIndex + >>> class IDim(DimensionIndex): ... >>> def field_op(inp: Field[[IDim], float64]): ... return inp >>> foast_tree = FieldOperatorParser.apply_to_function(field_op) diff --git a/src/gt4py/next/ffront/func_to_past.py b/src/gt4py/next/ffront/func_to_past.py index 8cf01eb691..e4722aa191 100644 --- a/src/gt4py/next/ffront/func_to_past.py +++ b/src/gt4py/next/ffront/func_to_past.py @@ -44,7 +44,7 @@ def func_to_past(inp: DSLProgramDef) -> PASTProgramDef: Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> @gtx.field_operator ... def copy(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 2ebdc8c74d..1b6cefc6b2 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -42,7 +42,7 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> stages.CompilableProgramDef: Example: >>> from gt4py import next as gtx >>> from gt4py.next.otf import arguments, workflow - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> @gtx.field_operator ... def copy(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: @@ -174,7 +174,8 @@ def _column_axis(all_closure_vars: dict[str, Any]) -> Optional[common.Dimension] if len(scanops_per_axis.values()) != 1: scanops_per_axis_str = "\n".join( - f"- {dim.__qualname__}: {', '.join(scanops)}" for dim, scanops in scanops_per_axis.items() + f"- {dim.__qualname__}: {', '.join(scanops)}" + for dim, scanops in scanops_per_axis.items() ) raise TypeError( @@ -246,7 +247,8 @@ class ProgramLowering( >>> from gt4py.next import Dimension, Field >>> >>> float64 = float - >>> IDim = Dimension("IDim") + >>> from gt4py.next.common import DimensionIndex + >>> class IDim(DimensionIndex): ... >>> >>> def fieldop(inp: Field[[IDim], "float64"]) -> Field[[IDim], "float64"]: ... >>> def program(inp: Field[[IDim], "float64"], out: Field[[IDim], "float64"]): diff --git a/src/gt4py/next/ffront/type_info.py b/src/gt4py/next/ffront/type_info.py index 900c0eeb05..c99a5fc3e0 100644 --- a/src/gt4py/next/ffront/type_info.py +++ b/src/gt4py/next/ffront/type_info.py @@ -8,7 +8,7 @@ import functools import inspect from collections.abc import Callable, Iterable -from typing import Final, Any, Iterator, Sequence, cast +from typing import Any, Final, Iterator, Sequence, cast import gt4py.next.ffront.type_specifications as ts_ffront import gt4py.next.type_system.type_specifications as ts @@ -188,13 +188,12 @@ def _scan_param_promotion( Example: -------- + >>> class I(common.DimensionIndex): ... >>> _scan_param_promotion( ... ts.ScalarType(kind=ts.ScalarKind.INT64), - ... ts.FieldType( - ... dims=[common.Dimension("I")], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64) - ... ), + ... ts.FieldType(dims=[I], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)), ... ) - FieldType(dims=[Dimension(value='I', kind=)], dtype=ScalarType(kind=, shape=None)) + FieldType(dims=[gt4py.next.ffront.type_info.I[horizontal]], dtype=ScalarType(kind=, shape=None)) """ def _as_field(dtype: ts.TypeSpec, path: tuple[int, ...]) -> ts.FieldType: diff --git a/src/gt4py/next/field_utils.py b/src/gt4py/next/field_utils.py index 7b9fe7e68b..4e2c6add71 100644 --- a/src/gt4py/next/field_utils.py +++ b/src/gt4py/next/field_utils.py @@ -36,10 +36,12 @@ def field_from_typespec( The tuple structure and dtype is taken from a type_specifications.DataType, which is either ScalarType or a CollectionTypeSpec of ScalarType (possibly nested). + >>> class I(common.DimensionIndex): ... >>> field_from_typespec( - ... ts.ScalarType(kind=ts.ScalarKind.INT32), common.domain({common.Dimension("I"): 1}), np + ... ts.ScalarType(kind=ts.ScalarKind.INT32), common.domain({I: 1}), np ... ) # doctest: +ELLIPSIS NumPyArrayField(... dtype=int32...) + >>> class I(common.DimensionIndex): ... >>> field_from_typespec( ... ts.TupleType( ... types=[ @@ -47,7 +49,7 @@ def field_from_typespec( ... ts.ScalarType(kind=ts.ScalarKind.FLOAT32), ... ] ... ), - ... common.domain({common.Dimension("I"): 1}), + ... common.domain({I: 1}), ... np, ... ) # doctest: +ELLIPSIS (NumPyArrayField(... dtype=int32...), NumPyArrayField(... dtype=float32...)) diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 766b6b2726..32f31c29a9 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -729,15 +729,15 @@ def _get_axes( In case all arguments are zero-dimensional return an empty sequence. >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.DimensionIndex): ... >>> i_field: LocatedField = _wrap_field( ... gtx.empty({IDim: range(3, 10)}, allocator=gtx.itir_python) ... ) >>> _get_axes((i_field, i_field)) - (Dimension(value='I', kind=),) + (gt4py.next.iterator.embedded.IDim[horizontal],) - >>> JDim = gtx.Dimension("J") + >>> class JDim(gtx.DimensionIndex): ... >>> j_field: LocatedField = _wrap_field( ... gtx.empty({JDim: range(3, 10)}, allocator=gtx.itir_python) ... ) @@ -753,7 +753,7 @@ def _get_axes( ValueError: Fields are defined on different axes. >>> _get_axes((i_field, zero_dim_field), ignore_zero_dims=True) - (Dimension(value='I', kind=),) + (gt4py.next.iterator.embedded.IDim[horizontal],) """ if isinstance(field_or_tuple, tuple): els_axes = [] diff --git a/src/gt4py/next/iterator/ir_utils/ir_makers.py b/src/gt4py/next/iterator/ir_utils/ir_makers.py index 27dbaa5c50..d6458f7a88 100644 --- a/src/gt4py/next/iterator/ir_utils/ir_makers.py +++ b/src/gt4py/next/iterator/ir_utils/ir_makers.py @@ -454,15 +454,15 @@ def domain( ranges_or_domain: dict[common.Dimension, tuple[itir.Expr, itir.Expr]] | common.Domain, ) -> itir.FunCall: """ - >>> IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL) - >>> JDim = common.Dimension(value="JDim", kind=common.DimensionKind.HORIZONTAL) + >>> class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + >>> class JDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... >>> str(domain(common.GridType.CARTESIAN, {IDim: (0, 10), JDim: (0, 20)})) - 'c⟨ IDimₕ: [0, 10[, JDimₕ: [0, 20[ ⟩' + 'c⟨ gt4py.next.iterator.ir_utils.ir_makers.IDimₕ: [0, 10[, gt4py.next.iterator.ir_utils.ir_makers.JDimₕ: [0, 20[ ⟩' >>> str(domain(common.GridType.UNSTRUCTURED, {IDim: (0, 10), JDim: (0, 20)})) - 'u⟨ IDimₕ: [0, 10[, JDimₕ: [0, 20[ ⟩' + 'u⟨ gt4py.next.iterator.ir_utils.ir_makers.IDimₕ: [0, 10[, gt4py.next.iterator.ir_utils.ir_makers.JDimₕ: [0, 20[ ⟩' >>> ij_domain = common.domain({IDim: (0, 10), JDim: (0, 20)}) >>> str(domain(common.GridType.UNSTRUCTURED, ij_domain)) - 'u⟨ IDimₕ: [0, 10[, JDimₕ: [0, 20[ ⟩' + 'u⟨ gt4py.next.iterator.ir_utils.ir_makers.IDimₕ: [0, 10[, gt4py.next.iterator.ir_utils.ir_makers.JDimₕ: [0, 20[ ⟩' """ if isinstance(ranges_or_domain, common.Domain): domain = ranges_or_domain @@ -592,9 +592,9 @@ def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: Examples -------- - >>> IDim = common.Dimension("IDim") + >>> class IDim(common.DimensionIndex): ... >>> str(broadcast("a", (IDim,))) - 'broadcast(a, {IDimₕ})' + 'broadcast(a, {gt4py.next.iterator.ir_utils.ir_makers.IDimₕ})' """ return call("broadcast")(expr, make_tuple(*(axis_literal(dim) for dim in dims))) diff --git a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py index f47d66cbe7..4bbc34e0f2 100644 --- a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py +++ b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py @@ -251,7 +251,11 @@ class FuseAsFieldOp( >>> from gt4py import next as gtx >>> from gt4py.next import utils >>> from gt4py.next.iterator.ir_utils import ir_makers as im - >>> IDim = gtx.Dimension("IDim") + >>> class IDim(gtx.DimensionIndex): ... + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # class declared in a doctest is not an attribute of the real module: + >>> import sys + >>> sys.modules[__name__].IDim = IDim >>> field_type = ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.INT32)) >>> d = im.domain("cartesian_domain", {IDim: (0, 1)}) >>> nested_as_fieldop = im.op_as_fieldop("plus", d)( @@ -261,8 +265,10 @@ class FuseAsFieldOp( ... im.ref("inp3", field_type), ... ) >>> print(nested_as_fieldop) - as_fieldop(λ(__arg0, __arg1) → ·__arg0 + ·__arg1, c⟨ IDimₕ: [0, 1[ ⟩)( - as_fieldop(λ(__arg0, __arg1) → ·__arg0 × ·__arg1, c⟨ IDimₕ: [0, 1[ ⟩)(inp1, inp2), inp3 + as_fieldop(λ(__arg0, __arg1) → ·__arg0 + ·__arg1, + c⟨ gt4py.next.iterator.transforms.fuse_as_fieldop.IDimₕ: [0, 1[ ⟩)( + as_fieldop(λ(__arg0, __arg1) → ·__arg0 × ·__arg1, + c⟨ gt4py.next.iterator.transforms.fuse_as_fieldop.IDimₕ: [0, 1[ ⟩)(inp1, inp2), inp3 ) >>> print( ... FuseAsFieldOp.apply( @@ -272,7 +278,8 @@ class FuseAsFieldOp( ... uids=utils.IDGeneratorPool(), ... ) ... ) - as_fieldop(λ(inp1, inp2, inp3) → ·inp1 × ·inp2 + ·inp3, c⟨ IDimₕ: [0, 1[ ⟩)(inp1, inp2, inp3) + as_fieldop(λ(inp1, inp2, inp3) → ·inp1 × ·inp2 + ·inp3, + c⟨ gt4py.next.iterator.transforms.fuse_as_fieldop.IDimₕ: [0, 1[ ⟩)(inp1, inp2, inp3) """ # noqa: RUF002 # ignore ambiguous multiplication character class Transformation(enum.Flag): diff --git a/src/gt4py/next/iterator/transforms/inline_fundefs.py b/src/gt4py/next/iterator/transforms/inline_fundefs.py index 2b8767e4a2..1158f7ed8e 100644 --- a/src/gt4py/next/iterator/transforms/inline_fundefs.py +++ b/src/gt4py/next/iterator/transforms/inline_fundefs.py @@ -44,7 +44,7 @@ def prune_unreferenced_fundefs(program: itir.Program) -> itir.Program: ... params=[im.sym("a")], ... expr=im.deref("a"), ... ) - >>> IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL) + >>> class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... >>> program = itir.Program( ... id="testee", ... function_definitions=[fun1, fun2], @@ -61,7 +61,7 @@ def prune_unreferenced_fundefs(program: itir.Program) -> itir.Program: >>> print(prune_unreferenced_fundefs(program)) testee(inp, out) { fun1 = λ(a) → ·a; - out @ c⟨ IDimₕ: [0, 10[ ⟩ ← fun1(inp); + out @ c⟨ gt4py.next.iterator.transforms.inline_fundefs.IDimₕ: [0, 10[ ⟩ ← fun1(inp); } """ fun_names = [fun.id for fun in program.function_definitions] 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 33d184e7e4..2bc8e84cc4 100644 --- a/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py +++ b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py @@ -84,7 +84,11 @@ class _PruneEmptyConcatWhere(PreserveLocationVisitor, NodeTranslator): `gt4py.next.iterator.transforms.concat_where.expand_tuple_args` before to prune them. >>> from gt4py.next import common - >>> IDim = common.Dimension("IDim") + >>> class IDim(common.DimensionIndex): ... + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # class declared in a doctest is not an attribute of the real module: + >>> import sys + >>> sys.modules[__name__].IDim = IDim >>> field_t = ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) >>> expr = im.concat_where( ... im.domain(common.GridType.CARTESIAN, {IDim: (10, itir.InfinityLiteral.POSITIVE)}), diff --git a/src/gt4py/next/iterator/transforms/remove_broadcast.py b/src/gt4py/next/iterator/transforms/remove_broadcast.py index 52c61105a4..db6445cc51 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -24,8 +24,13 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): Example: >>> from gt4py.next import Dimension, common - >>> IDim = Dimension("IDim") - >>> JDim = Dimension("JDim") + >>> from gt4py.next.common import DimensionIndex + >>> class IDim(DimensionIndex): ... + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # class declared in a doctest is not an attribute of the real module: + >>> class JDim(DimensionIndex): ... + >>> import sys + >>> sys.modules[__name__].IDim, sys.modules[__name__].JDim = IDim, JDim >>> domain = im.domain(common.GridType.CARTESIAN, {IDim: (0, 10), JDim: (0, 10)}) >>> expr = im.call("broadcast")( ... im.ref("inp"), @@ -36,7 +41,10 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): >>> expr.annex.domain = domain_utils.SymbolicDomain.from_expr(domain) >>> transformed = RemoveBroadcast.apply(expr) >>> print(transformed) - as_fieldop(deref, c⟨ IDimₕ: [0, 10[, JDimₕ: [0, 10[ ⟩)(inp) + as_fieldop( + deref, + c⟨ gt4py.next.iterator.transforms.remove_broadcast.IDimₕ: [0, 10[, gt4py.next.iterator.transforms.remove_broadcast.JDimₕ: [0, 10[ ⟩ + )(inp) """ PRESERVED_ANNEX_ATTRS = ("domain",) diff --git a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py index 868f9c81b6..26c54abf94 100644 --- a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py +++ b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py @@ -54,8 +54,8 @@ class ReplaceGetDomainRangeWithConstants(PreserveLocationVisitor, NodeTranslator Example: >>> from gt4py import next as gtx - >>> KDim = common.Dimension(value="KDim", kind=common.DimensionKind.VERTICAL) - >>> Vertex = common.Dimension(value="Vertex", kind=common.DimensionKind.HORIZONTAL) + >>> class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + >>> class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... >>> sizes = { ... "out": gtx.domain({Vertex: (0, 10), KDim: (0, 20)}), @@ -89,7 +89,8 @@ class ReplaceGetDomainRangeWithConstants(PreserveLocationVisitor, NodeTranslator >>> result = ReplaceGetDomainRangeWithConstants.apply(ir, sizes=sizes) >>> print(result) test(inp, out) { - out @ u⟨ Vertexₕ: [{0, 10}[0], {0, 10}[1][, KDimᵥ: [{0, 20}[0], {0, 20}[1][ ⟩ ← (⇑deref)(inp); + out @ u⟨ gt4py.next.iterator.transforms.replace_get_domain_range_with_constants.Vertexₕ: [{0, 10}[0], {0, 10}[1][, gt4py.next.iterator.transforms.replace_get_domain_range_with_constants.KDimᵥ: [{0, 20}[0], {0, 20}[1][ ⟩ + ← (⇑deref)(inp); } """ diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index a1622b0351..bcda84959e 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -405,15 +405,17 @@ def _canonicalize_nb_fields( Transform neighbor / sparse field type by removal of local dimension and addition of corresponding `ListType` dtype. Examples: + >>> class Vertex(common.DimensionIndex): ... + >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... >>> input_field = ts.FieldType( ... dims=[ - ... common.Dimension(value="Vertex"), - ... common.Dimension(value="V2E", kind=common.DimensionKind.LOCAL), + ... Vertex, + ... V2E, ... ], ... dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64), ... ) >>> _canonicalize_nb_fields(input_field) - FieldType(dims=[Dimension(value='Vertex', kind=)], dtype=ListType(element_type=ScalarType(kind=, shape=None), offset_type=Dimension(value='V2E', kind=))) + FieldType(dims=[gt4py.next.iterator.type_system.type_synthesizer.Vertex[horizontal]], dtype=ListType(element_type=ScalarType(kind=, shape=None), offset_type=gt4py.next.iterator.type_system.type_synthesizer.V2E[local])) """ match input_: case tuple() | ts.TupleType(): @@ -474,12 +476,12 @@ def _resolve_dimensions( tells you the dimensions of the field returned by the `as_fieldop`, in this case `[Vertex, K]`. - >>> Edge = common.Dimension(value="Edge") - >>> Vertex = common.Dimension(value="Vertex") - >>> Cell = common.Dimension(value="Cell") - >>> K = common.Dimension(value="K", kind=common.DimensionKind.VERTICAL) - >>> V2E = common.Dimension(value="V2E") - >>> C2V = common.Dimension(value="C2V") + >>> class Edge(common.DimensionIndex): ... + >>> class Vertex(common.DimensionIndex): ... + >>> class Cell(common.DimensionIndex): ... + >>> class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + >>> class V2E(common.DimensionIndex): ... + >>> class C2V(common.DimensionIndex): ... >>> input_dims = [Edge, K] >>> shift_tuple = ( ... itir.OffsetLiteral(value="C2V"), @@ -504,11 +506,22 @@ def _resolve_dimensions( ... ), ... } >>> _resolve_dimensions(input_dims, shift_tuple, offset_provider_type) - [Dimension(value='Cell', kind=), Dimension(value='K', kind=)] + [gt4py.next.iterator.type_system.type_synthesizer.Cell[horizontal], gt4py.next.iterator.type_system.type_synthesizer.K[vertical]] >>> from gt4py.next.iterator.ir_utils import ir_makers as im - >>> IDim = common.Dimension(value="IDim") + >>> class IDim(common.DimensionIndex): ... >>> IHalfDim = common.flip_staggered(IDim) - >>> JDim = common.Dimension(value="JDim") + >>> class JDim(common.DimensionIndex): ... + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # class declared in a doctest is not an attribute of the real module: + >>> import sys + >>> sys.modules[__name__].Edge = Edge + >>> sys.modules[__name__].Vertex = Vertex + >>> sys.modules[__name__].Cell = Cell + >>> sys.modules[__name__].K = K + >>> sys.modules[__name__].V2E = V2E + >>> sys.modules[__name__].C2V = C2V + >>> sys.modules[__name__].IDim = IDim + >>> sys.modules[__name__].JDim = JDim >>> JHalfDim = common.flip_staggered(JDim) >>> input_dims = [IDim, JDim] >>> shift_tuple = ( @@ -524,7 +537,7 @@ def _resolve_dimensions( ... itir.OffsetLiteral(value=0), ... ) >>> _resolve_dimensions(input_dims, shift_tuple, offset_provider_type) - [Dimension(value='JDim', kind=), Dimension(value='IDim', kind=)] + [gt4py.next.iterator.type_system.type_synthesizer.JDim[horizontal], gt4py.next.iterator.type_system.type_synthesizer.IDim[horizontal]] """ resolved_dims = [] diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 0527d019df..b27625c4f5 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -102,7 +102,7 @@ def primitive_constituents( Return the primitive types contained in a composite type. >>> from gt4py.next import common - >>> I = common.Dimension(value="I") + >>> class I(common.DimensionIndex): ... >>> int_type = ts.ScalarType(kind=ts.ScalarKind.INT64) >>> field_type = ts.FieldType(dims=[I], dtype=int_type) @@ -391,10 +391,10 @@ def extract_dims(symbol_type: ts.TypeSpec) -> list[common.Dimension]: Examples: >>> extract_dims(ts.ScalarType(kind=ts.ScalarKind.INT64, shape=[3, 4])) [] - >>> I = common.Dimension(value="I") - >>> J = common.Dimension(value="J") + >>> class I(common.DimensionIndex): ... + >>> class J(common.DimensionIndex): ... >>> extract_dims(ts.FieldType(dims=[I, J], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64))) - [Dimension(value='I', kind=), Dimension(value='J', kind=)] + [gt4py.next.type_system.type_info.I[horizontal], gt4py.next.type_system.type_info.J[horizontal]] """ if isinstance(symbol_type, ts.ScalarType): return [] @@ -408,8 +408,8 @@ def is_local_field(type_: ts.FieldType) -> bool: Return if `type_` is a field defined on a local dimension. Examples: - >>> V = common.Dimension(value="V") - >>> V2E = common.Dimension(value="V2E", kind=common.DimensionKind.LOCAL) + >>> class V(common.DimensionIndex): ... + >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... >>> is_local_field( ... ts.FieldType(dims=[V, V2E], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64)) ... ) @@ -442,7 +442,7 @@ def is_compatible_type(type_a: ts.TypeSpec, type_b: ts.TypeSpec) -> bool: Beside that this function simply checks for equality of types. >>> bool_type = ts.ScalarType(kind=ts.ScalarKind.BOOL) - >>> IDim = common.Dimension(value="IDim") + >>> class IDim(common.DimensionIndex): ... >>> type_on_i_of_i_it = it_ts.IteratorType( ... position_dims=[IDim], defined_dims=[IDim], element_type=bool_type ... ) @@ -452,7 +452,7 @@ def is_compatible_type(type_a: ts.TypeSpec, type_b: ts.TypeSpec) -> bool: >>> is_compatible_type(type_on_i_of_i_it, type_on_undefined_of_i_it) True - >>> JDim = common.Dimension(value="JDim") + >>> class JDim(common.DimensionIndex): ... >>> type_on_j_of_j_it = it_ts.IteratorType( ... position_dims=[JDim], defined_dims=[JDim], element_type=bool_type ... ) @@ -570,7 +570,11 @@ def promote( `offset_type` of `None` (a list from `make_const_list`) is compatible with any other. >>> dtype = ts.ScalarType(kind=ts.ScalarKind.INT64) - >>> I, J, K = (common.Dimension(value=dim) for dim in ["I", "J", "K"]) + >>> class I(common.DimensionIndex): ... + + >>> class J(common.DimensionIndex): ... + + >>> class K(common.DimensionIndex): ... >>> promoted: ts.FieldType = promote( ... ts.FieldType(dims=[I, J], dtype=dtype), ts.FieldType(dims=[I, J, K], dtype=dtype), dtype ... ) @@ -583,7 +587,7 @@ def promote( >>> promoted.dims == [I, J, K] and promoted.dtype == dtype True - >>> V2E = common.Dimension(value="V2E", kind=common.DimensionKind.LOCAL) + >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... >>> list_dtype = ts.ListType(element_type=dtype, offset_type=V2E) >>> promote( ... ts.FieldType(dims=[I], dtype=list_dtype), From 3109d549ab69d483039ea0cbe7fe93b98dc71900 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 16:00:16 +0200 Subject: [PATCH 10/20] test[next]: migrate typing_tests to class-based dimensions With the dimension half of the mypy plugin gone, a dimension bound to a variable (`IDim = gtx.Dimension("I")`) is no longer valid in an annotation: only a class is. `test_typing_exports` runs these snippets in CI and was not covered by any local suite. Taken from #2844 unchanged. Static typing depends only on dimensions being classes, which both designs share; nothing in these snippets touches identity semantics, and main has not changed the file since #2844's base. 16 passed. --- typing_tests/test_next.yaml | 38 ++++++++++++++++++------------------- 1 file changed, 19 insertions(+), 19 deletions(-) diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 5db040f8df..20ed020333 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -4,7 +4,7 @@ import typing from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -16,7 +16,7 @@ main: | from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator(grid_type=None, backend=None) def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -26,7 +26,7 @@ main: | from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator(grid_type=None) def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -36,7 +36,7 @@ main: | from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.scan_operator(axis=KDim) def somescanop( @@ -49,7 +49,7 @@ main: | from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.scan_operator(axis=KDim, forward=True, init=0.0) def somescanop( @@ -64,7 +64,7 @@ import typing from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -85,7 +85,7 @@ import typing from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -106,7 +106,7 @@ import typing from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -127,8 +127,8 @@ import typing from gt4py import next as gtx - CellDim = gtx.Dimension("CellDim", kind=gtx.DimensionKind.HORIZONTAL) - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class CellDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... T = typing.TypeVar("T", float, gtx.float32, gtx.float64, bool, gtx.int32, gtx.int64) CellKField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim, KDim], T] @@ -145,8 +145,8 @@ import typing from gt4py import next as gtx - CellDim = gtx.Dimension("CellDim", kind=gtx.DimensionKind.HORIZONTAL) - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class CellDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... T = typing.TypeVar("T", float, gtx.float32, gtx.float64, bool, gtx.int32, gtx.int64) CellKField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim, KDim], T] CKTuple: typing.TypeAlias = tuple[CellKField[T], CellKField[T]] @@ -163,8 +163,8 @@ from gt4py import next as gtx from gt4py.next import where - CellDim = gtx.Dimension("CellDim", kind=gtx.DimensionKind.HORIZONTAL) - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class CellDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... T = typing.TypeVar("T", float, gtx.float32, gtx.float64, bool, gtx.int32, gtx.int64) CellField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim], T] CellKField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim, KDim], T] @@ -205,7 +205,7 @@ import typing from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -218,7 +218,7 @@ main: | from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.scan_operator(axis=KDim) def somescanop( @@ -236,7 +236,7 @@ import typing from gt4py import next as gtx - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -260,8 +260,8 @@ import typing from gt4py import next as gtx - CDim = gtx.Dimension("CDim", kind=gtx.DimensionKind.HORIZONTAL) - KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + class CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo( From ce595075f9dfe1d126ad2b0e0432269aea785c56 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 21 Sep 2026 22:40:58 +0200 Subject: [PATCH 11/20] style[next]: ruff format lines shortened by the tag migration `pre-commit run --all-files` (what CI runs) reformats four lines that now fit on one line after `.value` -> `.tag`; the per-commit runs only covered staged files. --- src/gt4py/next/iterator/ir_utils/domain_utils.py | 4 +--- .../runners/dace/lowering/gtir_to_sdfg_primitives.py | 4 +--- .../runners/dace/transformations/loop_blocking.py | 4 +++- src/gt4py/next/type_system/type_specifications.py | 6 +++++- 4 files changed, 10 insertions(+), 8 deletions(-) diff --git a/src/gt4py/next/iterator/ir_utils/domain_utils.py b/src/gt4py/next/iterator/ir_utils/domain_utils.py index 80330e8cc8..b23ef3a934 100644 --- a/src/gt4py/next/iterator/ir_utils/domain_utils.py +++ b/src/gt4py/next/iterator/ir_utils/domain_utils.py @@ -155,9 +155,7 @@ def from_expr(cls, node: itir.Node) -> SymbolicDomain: axis_literal, lower_bound, upper_bound = named_range.args assert isinstance(axis_literal, itir.AxisLiteral) - ranges[common.resolve(axis_literal.value)] = ( - SymbolicRange(lower_bound, upper_bound) - ) + ranges[common.resolve(axis_literal.value)] = SymbolicRange(lower_bound, upper_bound) return cls(_GRID_TYPE_MAPPING[node.fun.id], ranges) def as_expr(self) -> itir.FunCall: 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 ee4d56f320..e265377a45 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -325,9 +325,7 @@ def _construct_if_branch_output( assert out_type.dtype.offset_type is not None assert isinstance(out_type.dtype.element_type, ts.ScalarType) dtype = gtx_dace_args.as_dace_type(out_type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type( - out_type.dtype.offset_type.tag - ) + offset_provider_type = sdfg_builder.get_offset_provider_type(out_type.dtype.offset_type.tag) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) shape = [*shape, offset_provider_type.max_neighbors] diff --git a/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py b/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py index 47b644a2b8..2e0d21eca9 100644 --- a/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py +++ b/src/gt4py/next/program_processors/runners/dace/transformations/loop_blocking.py @@ -111,7 +111,9 @@ def __init__( super().__init__() if blocking_parameters is not None: self.blocking_parameters = [ - gtx_dace_lowering.get_map_variable(p) if isinstance(p, gtx_common.DimensionMeta) else p + gtx_dace_lowering.get_map_variable(p) + if isinstance(p, gtx_common.DimensionMeta) + else p for p in blocking_parameters ] else: diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index a2bfc86831..806350542e 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -122,7 +122,11 @@ class FieldType(DataType, CallableType): dtype: ScalarType | ListType def __str__(self) -> str: - dims = "..." if self.dims is Ellipsis else f"[{', '.join(dim.__qualname__ for dim in self.dims)}]" + dims = ( + "..." + if self.dims is Ellipsis + else f"[{', '.join(dim.__qualname__ for dim in self.dims)}]" + ) return f"Field[{dims}, {self.dtype}]" @eve_datamodels.validator("dims") From 0e55c87982f486e33681c49027c3642c04650140 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 15:28:55 +0200 Subject: [PATCH 12/20] fix[next]: review fixes for dimension classes - a dimension cannot be staggered twice: the check reads the *argument* - intern staggered dimensions with 'setdefault', since compilation runs in threads - 'resolve' rejects a bracketed tag whose owner is not a parametrized dimension - the 'Dim + offset' hint quotes the dimension's display name, not '.value' - ADR 0026: the name-prefix encoding is superseded by ADR 0028 - drop the FieldOffset tag/name analysis: the knowledge base keeps that record - 'runtimes': say where an unpicklable job is actually reported --- .../ADRs/next/0026-Staggered_Dimensions.md | 20 +- .../next/fieldoffset-tag-constraints.md | 551 ------------------ src/gt4py/next/common.py | 19 +- .../ffront/foast_passes/type_deduction.py | 4 +- src/gt4py/next/otf/runners.py | 4 +- tests/next_tests/unit_tests/test_common.py | 17 + 6 files changed, 52 insertions(+), 563 deletions(-) delete mode 100644 docs/development/next/fieldoffset-tag-constraints.md diff --git a/docs/development/ADRs/next/0026-Staggered_Dimensions.md b/docs/development/ADRs/next/0026-Staggered_Dimensions.md index c8fbd3a91f..9cde0a7179 100644 --- a/docs/development/ADRs/next/0026-Staggered_Dimensions.md +++ b/docs/development/ADRs/next/0026-Staggered_Dimensions.md @@ -7,7 +7,13 @@ tags: [] - **Status**: valid - **Authors**: Till Ehrengruber (@tehrengruber) - **Created**: 2026-07-08 -- **Updated**: 2026-07-09 +- **Updated**: 2026-09-23 + +> The *encoding* of this record is superseded by +> [ADR 0028](0028-Dimensions_As_Nominal_Types.md): a staggered dimension is the +> real, interned class `Staggered[D]`, not a name prefix. The semantics below -- +> half-integer positions, the shift convention, and the gtfn and DaCe treatment +> -- are unchanged. A **staggered dimension** is a dimension sitting at the **half-integer** positions of a base dimension. For example, in a cell-centered 2D Cartesian grid @@ -58,10 +64,12 @@ index arithmetic is encoded in `common.connectivity_for_cartesian_shift`. ## Encoding -A staggered dimension is encoded as its base dimension's name with the internal -`_Staggered` prefix (`common._STAGGERED_PREFIX`), rather than as a new attribute -on `Dimension`. The helpers `is_staggered`, `flip_staggered` and -`as_non_staggered` operate purely on that prefix. +A staggered dimension was encoded as its base dimension's name with the internal +`_Staggered` prefix, rather than as a new attribute on `Dimension`. Since +[ADR 0028](0028-Dimensions_As_Nominal_Types.md) it is the class +`common.Staggered[D]`, whose identity is the class itself; the helpers +`is_staggered`, `flip_staggered` and `as_non_staggered` are unchanged in meaning +and now read the class's `base`. Dimensions are identified by their **name** and appear throughout the toolchain in more than one form: as a `common.Dimension` instance, but also as an @@ -87,7 +95,7 @@ by the `as_non_staggered` name. In gtfn a staggered dimension is emitted as a C++ `using` alias of its base dimension's tag (`_add_staggered_aliases` in `itir_to_gtfn_ir.py`): e.g. -`_StaggeredIDim_t` becomes an alias of `IDim_t`, and `visit_AxisLiteral` +the staggered tag's mangled name becomes an alias of `IDim_t`, and `visit_AxisLiteral` correspondingly emits the base dimension name. The reason is that gtfn lowers a shift to an integer offset along a SID axis, and diff --git a/docs/development/next/fieldoffset-tag-constraints.md b/docs/development/next/fieldoffset-tag-constraints.md deleted file mode 100644 index 45122c5a49..0000000000 --- a/docs/development/next/fieldoffset-tag-constraints.md +++ /dev/null @@ -1,551 +0,0 @@ -# Tag and name constraints for `FieldOffset`, `Dimension` and offset providers - -Analysis of the *string-identity* assumptions that connect `FieldOffset` tags, -`Dimension` names and `offset_provider` keys in `gt4py.next`, across the -embedded, IR and backend contexts. - -Status: descriptive — this documents the implementation as it is, it does not -propose a change. Line references are against `main` at `b3c53fa7e` -(v1.2.2, 2026-09-03). - -Related: [ADR 0019 Connectivities](../ADRs/next/0019-Connectivities.md), -[ADR 0026 Staggered Dimensions](../ADRs/next/0026-Staggered_Dimensions.md). - -## 0. The five name spaces - -| # | Name space | Type | Defined at | -| ------ | ------------------------------------------------------------------ | ---------------------------------------------------- | ---------------------------------------------------------- | -| **N1** | `FieldOffset.value` — the *offset tag* | `str` (from `runtime.Offset.value: Union[int, str]`) | `iterator/runtime.py:36-37`, `ffront/fbuiltins.py:471-472` | -| **N2** | The **Python closure-variable name** the `FieldOffset` is bound to | `str` | consumed at `ffront/foast_to_gtir.py:305, 331` | -| **N3** | `Dimension.value` — the *dimension tag* | `str` | `common.py:79-81` | -| **N4** | `offset_provider` **dict key** | `str` | `common.py:1200-1209` | -| **N5** | ITIR `OffsetLiteral.value` / `AxisLiteral.value` | `str` | `iterator/ir.py:88-96` | - -`ts.OffsetType` carries **only** `source`/`target` -(`type_system/type_specifications.py:73-76`); N1 is discarded at -`fbuiltins.py:484-485`. That erasure is the root cause of most rows below. - -The only validation `FieldOffset` performs on itself is on the *kind*, never on -any name (`ffront/fbuiltins.py:470-482`): - -```python -class FieldOffset(runtime.Offset): # .value is the tag, inherited from runtime.Offset - source: common.Dimension - target: tuple[Dimension] | tuple[Dimension, Dimension] - - def __post_init__(self) -> None: - if len(self.target) == 2 and self.target[1].kind != common.DimensionKind.LOCAL: - raise ValueError("Second dimension in offset must be a local dimension.") -``` - -## 1. Concept inventory - -Every class and type alias in `src/gt4py/next/` that participates in the -offset/connectivity vocabulary, grouped by layer. - -### L0 — Vocabulary (`common.py`) - -| Concept | Where | Role | -| ------------------------------------------------- | ------------------------------ | -------------------------------------------------------------------------- | -| `Tag = str` | `common.py:63` | the alias that makes every name space stringly-typed | -| `DimensionKind` | `common.py:66-72` | `HORIZONTAL` / `VERTICAL` / `LOCAL` | -| `Dimension` | `common.py:79-81` | `(value: str, kind)`; `__add__`/`__sub__` build Cartesian shifts | -| `UnitRange`, `NamedRange`, `NamedIndex`, `Domain` | `common.py:196, 358, 369, 432` | index-space vocabulary; `Domain.dims` is where local dims appear on fields | - -### L1 — Connectivity objects (runtime data) - -| Concept | Where | Role | -| ----------------------------- | ----------------------------- | ------------------------------------------------------------------- | -| `Connectivity` | `common.py:990` | `Protocol`; a `Field` of indices with a `codomain` | -| `GatherConnectivity` | `common.py:1099` | nominal (not a `Protocol`): `premap` is a data-moving gather | -| `NeighborTable` | `common.py:1149` | `Protocol`; 2-D table-backed neighbor connectivity | -| `NdArrayConnectivityField` | `nd_array_field.py:516` | the concrete implementation | -| `NumPyArrayConnectivityField` | `nd_array_field.py:1032` | array-library variant | -| `CuPyArrayConnectivityField` | `nd_array_field.py:1049` | array-library variant | -| `JaxArrayConnectivityField` | `nd_array_field.py:1087` | array-library variant | -| `CartesianConnectivity` | `common.py:1241` | affine shift; no `ndarray`, **not** a `GatherConnectivity` | -| `StridedConnectivityField` | `iterator/embedded.py:107` | incomplete; iterator-view only (`TODO(havogt)`) | -| `_ConnectivityFileRef` | `otf/compilation_tasks.py:50` | lazy pickling stand-in; dumps to `.npy` to cross process boundaries | - -Constructors: `constructors.as_connectivity`, plus the `_field` / `_connectivity` -`singledispatch` pair at `common.py:1121-1146`. - -### L2 — Connectivity *types* (compile time) - -| Concept | Where | Contents | -| -------------------------- | ------------------- | --------------------------------------------------------------------------- | -| `ConnectivityType` | `common.py:964-973` | `domain`, `codomain`, `skip_value`, `dtype` | -| `NeighborConnectivityType` | `common.py:976-986` | `+ max_neighbors`; `.source_dim == domain[0]`, `.neighbor_dim == domain[1]` | - -### L3 — The provider (name to data binding) - -| Concept | Where | -| ------------------------------------------------------------------------ | --------------------- | -| `OffsetProvider = Mapping[Tag, NeighborTable]` | `common.py:1177` | -| `OffsetProviderType = Mapping[Tag, NeighborConnectivityType]` | `common.py:1178` | -| `OffsetProviderElem`, `OffsetProviderTypeElem` | `common.py:1172-1173` | -| `get_offset`, `get_offset_type`, `has_offset`, `offset_provider_to_type` | `common.py:1193-1221` | - -### L4 — Frontend declarations - -| Concept | Where | Contents | -| ---------------------------------- | --------------------------- | ---------------------------------------------- | -| `runtime.Offset` | `iterator/runtime.py:36-37` | `value: int \| str` | -| `FieldOffset(runtime.Offset)` | `ffront/fbuiltins.py:472` | `+ source`, `target` | -| `as_offset` builtin | `ffront/experimental.py:17` | dynamic Cartesian shift from an index field | -| `connectivity_for_cartesian_shift` | `common.py:1467` | builds a `CartesianConnectivity`; needs no tag | - -### L5 — Frontend types - -| Concept | Where | Note | -| ------------------- | ------------------------------ | ---------------------------------------------------- | -| `ts.OffsetType` | `type_specifications.py:73-79` | `source`/`target` only — **the tag is dropped here** | -| `ts.DimensionType` | `type_specifications.py:55` | wraps a `Dimension` | -| `ts.FieldType.dims` | `type_specifications.py:121` | sparse fields carry the local dim as a list member | - -### L6 — ITIR nodes - -| Node | Where | Carries | -| ---------------------- | ----------------------- | --------------------------------------------------- | -| `itir.OffsetLiteral` | `iterator/ir.py:88-89` | a bare `str` tag — unstructured | -| `itir.AxisLiteral` | `iterator/ir.py:92-96` | `value` + `kind` — a serialized `Dimension` | -| `itir.CartesianOffset` | `iterator/ir.py:99-101` | two `AxisLiteral`s — **no tag, no provider lookup** | - -### L7 — ITIR types - -| Concept | Where | Contents | -| --------------------------- | ------------------------------------------------ | ------------------------------------------------- | -| `it_ts.OffsetLiteralType` | `iterator/type_system/type_specifications.py:19` | `value: ScalarType \| str` | -| `it_ts.CartesianOffsetType` | `…:23` | `domain`, `codomain` | -| `it_ts.NamedRangeType` | `…:15` | `dim` | -| `it_ts.IteratorType` | `…:28` | `position_dims`, `defined_dims` | -| `ts.ListType` | `type_specifications.py:108-118` | `element_type` + `offset_type: Dimension \| None` | - -`ListType`'s docstring states the frontend/IR split explicitly: *"not used in the -frontend. The concept is represented as Field with local Dimension."* - -### L8 — Embedded iterator runtime - -| Concept | Where | Role | -| ---------------------------------- | --------------------------------- | --------------------------------------------------------- | -| `SparseTag(Tag)` | `iterator/embedded.py:102` | marks a shift into the sparse axis | -| `MDIterator`, `SparseListIterator` | `iterator/embedded.py:~800, 1507` | iterators; the latter holds `list_offset: Tag` | -| `_List`, `_ConstList` | `iterator/embedded.py:1399, 1420` | neighbor-list values | -| `_CONST_DIM` | `iterator/embedded.py:220` | reserved LOCAL dim, deliberately absent from the provider | -| position dicts | `iterator/embedded.py:597-616` | keyed by `Dimension.value` strings | - -### L9 — Backend representations - -| Concept | Where | Role | -| ------------------------------------------- | ------------------------------------------- | ----------------------------------------------------- | -| `gtfn_ir.OffsetLiteral` | `gtfn/gtfn_ir.py:52` | lowered tag | -| `gtfn_ir.TagDefinition` | `gtfn/gtfn_ir.py:249-251` | `name`, optional `alias`; emits `generated::_t` | -| `gtfn_ir.UnstructuredDomain.connectivities` | `gtfn/gtfn_ir.py:90-93` | `SymRef` to an offset declaration | -| `gtfn_ir.TaggedValues` | `gtfn/gtfn_ir.py:80-82` | tag-keyed sizes/offsets | -| `dace FieldopData` | `dace/lowering/gtir_to_sdfg_types.py:27-34` | carries the local-dim/offset-provider association | -| `dace connectivity_identifier` | `dace/sdfg_args.py:56-70` | `gt_conn_` array naming | - -### Concept count - -| Kind | Count | Notes | -| ------------------------------------------ | ----- | ---------------------------------------------------------------------------------------------------------------------- | -| Runtime connectivity classes | 8 | 3 are array-library variants of one; 1 is incomplete | -| Connectivity type classes | 2 | | -| Declaration classes | 2 | the subclassing is flagged as a conceptual mismatch at `fbuiltins.py:467` | -| Type-system representations of "an offset" | 5 | `ts.OffsetType`, `it_ts.OffsetLiteralType`, `it_ts.CartesianOffsetType`, `ts.ListType.offset_type`, `ts.DimensionType` | -| IR node kinds | 4 | 3 ITIR + 1 GTFN | -| Provider aliases | 4 | | - -Roughly **25 distinct concepts** for what is conceptually one thing — a mapping -between two index spaces — plus a name for it. - -## 2. How the concepts relate - -### 2.1 Connectivity class hierarchy - -```text -Field (Protocol) -└── Connectivity (Protocol) common.py:990 - ├── GatherConnectivity <- nominal, gather premap common.py:1099 - │ └── NeighborTable (Protocol, 2-D, table-backed) common.py:1149 - │ └── NdArrayConnectivityField nd_array_field.py:516 - │ ├── NumPyArrayConnectivityField nd_array_field.py:1032 - │ ├── CuPyArrayConnectivityField nd_array_field.py:1049 - │ └── JaxArrayConnectivityField nd_array_field.py:1087 - ├── CartesianConnectivity <- affine, no ndarray common.py:1241 - └── StridedConnectivityField <- WIP, iterator only iterator/embedded.py:107 -``` - -```mermaid -classDiagram - class Field { - <> - } - class Connectivity { - <> - +codomain: Dimension - +__gt_type__() ConnectivityType - } - class GatherConnectivity { - +ndarray - } - class NeighborTable { - <> - +__gt_type__() NeighborConnectivityType - } - class CartesianConnectivity { - +domain_dim - +offset: int - } - class StridedConnectivityField - class ConnectivityType { - +domain: tuple~Dimension~ - +codomain: Dimension - +skip_value - +dtype - } - class NeighborConnectivityType { - +max_neighbors: int - +source_dim - +neighbor_dim - } - Field <|-- Connectivity - Connectivity <|-- GatherConnectivity - Connectivity <|-- CartesianConnectivity - Connectivity <|-- StridedConnectivityField - GatherConnectivity <|-- NeighborTable - NeighborTable <|-- NdArrayConnectivityField - NdArrayConnectivityField <|-- NumPyArrayConnectivityField - NdArrayConnectivityField <|-- CuPyArrayConnectivityField - NdArrayConnectivityField <|-- JaxArrayConnectivityField - ConnectivityType <|-- NeighborConnectivityType - Connectivity ..> ConnectivityType : __gt_type__() - NeighborTable ..> NeighborConnectivityType : __gt_type__() -``` - -### 2.2 Declaration vs type vs data — the duplicated triple - -`FieldOffset` carries exactly the information in `NeighborConnectivityType` plus -a name, with **inverted vocabulary** and no cross-check. `FieldOffset.source` is -the connectivity's *codomain*; `FieldOffset.target` is its *domain*. The -inversion is because `source`/`target` describe the *field remap* (the field -lives on `source` and ends up on `target`), while `domain`/`codomain` describe -the *table*. - -```text - DECLARATION TYPE DATA - ─────────── ──── ──── - FieldOffset ts.OffsetType (none — bound later) - .value ─── dropped ──X - .target[0] ═══════════════ .target[0] ═══ A8 ═══════════ ConnectivityType.domain[0] - .target[1] ═══════════════ .target[1] ═══ A6 ═══════════ ConnectivityType.domain[1] - .source ═══════════════ .source ═══ A7 ═══════════ ConnectivityType.codomain - ^^^^^^^^^^^^^^^^^^^^^^^^^ - the same information, authored twice, - with inverted vocabulary, never cross-checked -``` - -### 2.3 Name flow — where the five name spaces diverge - -Four independently-authored strings converge on one dict lookup, and *which* of -them arrives there depends on the execution path and the operation. - -```text - ┌──────────────────────────────────────────────────┐ - │ V2EDim = Dimension("V2E", LOCAL) (N3) │ - USER AUTHORS │ V2E = FieldOffset("V2E", Edge,(V,V2EDim)) │ - FOUR STRINGS │ ^^^ (N2) │ - │ ^^^^^ (N1) │ - │ offset_provider = {"V2E": table} (N4) │ - └──────────────────────────────────────────────────┘ - │ - ┌─────────────────────────┴─────────────────────────┐ - │ │ - EMBEDDED PATH COMPILED PATH - │ │ - ┌──────────┴──────────┐ ┌─────────────┴─────────────┐ - │ shift │ reduce │ FOAST -> GTIR │ - │ fbuiltins.py:494 │ nd_array_field.py:983 │ foast_to_gtir.py:305,331 │ - │ uses N1 │ uses N3 (axis.value) │ uses N2 (Name.id) -> N5 │ - └──────────┬──────────┘ └─────────────┬─────────────┘ - │ │ - │ ┌──────────────────────────┤ - │ │ reduce: unroll_reduce.py:47 - │ │ uses N3 (ListType.offset_type.value) - │ │ - │ │ sparse arg: gtfn_module.py:95 - │ │ gtir_to_sdfg.py:581 - │ │ uses N3 (dim.value) - │ │ - └────────────┬───────────┴─────────────┬────────────┘ - v v - get_offset(offset_provider, ) == N4 - │ - v - NeighborTable / NeighborConnectivityType -``` - -```mermaid -flowchart TD - subgraph AUTHORED["User authors four strings"] - N3["N3 - Dimension('V2E', LOCAL)"] - N1["N1 - FieldOffset.value = 'V2E'"] - N2["N2 - python variable name V2E"] - N4["N4 - offset_provider key 'V2E'"] - end - - N1 --> EshiftE["embedded shift
fbuiltins.py:494"] - N3 --> EredE["embedded reduce
nd_array_field.py:983"] - N2 --> LOW["FOAST to GTIR
foast_to_gtir.py:305, 331"] - LOW --> N5["N5 - itir.OffsetLiteral"] - N5 --> CshiftC["compiled shift
type_synthesizer.py:748"] - N3 --> CredC["compiled reduce
unroll_reduce.py:47"] - N3 --> SPARSE["sparse field argument
gtfn_module.py:95
gtir_to_sdfg.py:581"] - - EshiftE --> GET - EredE --> GET - CshiftC --> GET - CredC --> GET - SPARSE --> GET - N4 -.->|"must equal the string that arrives"| GET - - GET["get_offset(offset_provider, string)
common.py:1200"] - GET --> DATA["NeighborTable / NeighborConnectivityType"] -``` - -### 2.4 The contrast that suggests the fix - -Cartesian shifts carry **dimensions** in the IR node; unstructured shifts carry -a **string** that must be resolved against a dict. Every constraint A1-A5 exists -only on the right-hand side. - -```text - CARTESIAN (already clean) UNSTRUCTURED (entangled) - ───────────────────────── ──────────────────────── - field(IDim + 1) field(V2E) - │ │ - v v - CartesianConnectivity itir.OffsetLiteral("V2E") <- a string - (common.py:1241) │ - │ v - v get_offset(provider, "V2E") - itir.CartesianOffset │ - domain: AxisLiteral v - codomain: AxisLiteral NeighborTable - (iterator/ir.py:99) - │ - v - NO tag. NO provider entry. NO lookup. -``` - -This is the concrete precedent behind any consolidation proposal: the Cartesian -path already eliminated the string indirection, and the unstructured path -retains it only because the neighbor table data must be supplied at runtime. - -## 3. Master table — cross-name-space identity constraints - -| # | Constraint | Embedded (field) | Embedded (iterator) | IR / type system | GTFN | DaCe | Enforced? | Source | -| ------- | ---------------------------------------------------------------------------- | ---------------- | ------------------- | ---------------- | ---------------- | ------------ | ----------------------------- | -------------------------------------------------------------------------------------------- | -| **A1** | `FieldOffset.value` (N1) == provider key (N4) | required | required | — | — | — | `KeyError` | `fbuiltins.py:494, 508`; `common.py:1207-1208` | -| **A2** | Python var name (N2) == provider key (N4) | — | — | required | required | required | silent; `KeyError` at runtime | `foast_to_gtir.py:305, 331` | -| **A3** | local dim `.value` (N3) == provider key (N4), **reductions** | required | required | required | required | required | `KeyError` | `nd_array_field.py:981-985`; `embedded.py:953, 1517, 1776`; `unroll_reduce.py:43-50, 61-65` | -| **A4** | local dim `.value` (N3) == provider key (N4), **sparse field args** | — | — | — | required | required | `assert` / `ValueError` | `gtfn_module.py:88-98`; `gtir_to_sdfg.py:572-585, 838-842`; `gtir_to_sdfg_lambda.py:766-770` | -| **A5** | `FieldOffset.value` (N1) == local dim `.value` (N3), **shift path** | n/a (A1 governs) | n/a | **not** required | **not** required | inconsistent | codegen branch handles it | `itir_to_gtfn_ir.py:181-190`; regression test | -| **A6** | `target[-1]` == connectivity `neighbor_dim` (full `Dimension` equality) | required | required | required | required | required | no eager check; index error | `fbuiltins.py:496`; `common.py:984-986` | -| **A7** | `FieldOffset.source` == connectivity `codomain` | required | required | required | required | required | `assert` only | `embedded.py:596-614`; `type_synthesizer.py:748-758` | -| **A8** | `FieldOffset.target[0]` == connectivity `domain[0]` (`source_dim`) | required | required | required | required | required | `assert found` | `type_synthesizer.py:752-758`; `embedded.py:597-599` | -| **A9** | `Dimension.value` (N3) is the key of the embedded **iterator position dict** | — | required | — | — | — | `assert ... in pos` | `embedded.py:574-576, 597-616, 941-950` | -| **A10** | `Dimension.value` (N3) round-trips through `AxisLiteral.value` (N5) | — | — | required | required | required | structural | `iterator/ir.py:92-96`; `ir_utils/misc.py:234-235`; `inference.py:463-464` | - -### Notes on A5 - -A5 is the only row with history. PR #1789 (`fix[next]: gtfn with offset name != local dimension name`) lifted it for shifts and added the -`if offset_name != connectivity_type.neighbor_dim.value` branch at -`itir_to_gtfn_ir.py:185-190`. Its regression test is -`tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py`, -whose docstring gives the motivation: - -> If the value of the `NeighborConnectivityType.neighbor_dim` did not match the -> `FieldOffset` value, gtfn would silently ignore the neighbor index, see -> . - -That test covers **only** `a(Off[1])` on `GTFN_CPU`. It does not cover -`neighbor_sum`, embedded execution, or DaCe. A3 and A4 were never lifted, so a -mismatch still breaks reductions and sparse arguments. - -The DaCe "inconsistent" entry: `gtir_to_sdfg_lambda.py:1155` builds -`Dimension(offset, LOCAL)` — a local dim named after the *tag* — while -`type_synthesizer.py:327-329` builds the same `ListType` from -`conn_type.neighbor_dim`. The two agree only when A5 holds. - -## 4. Per-context detail - -### 4.1 Embedded — field level (`nd_array_field`, `fbuiltins`) - -| Site | Key used | Constraint | -| ------------------------------------------------ | ----------------- | ------------------------------------------------------------------------------ | -| `fbuiltins.py:491-498` `FieldOffset.__getitem__` | `self.value` (N1) | A1; then `NamedIndex(self.target[-1], offset)` gives A6 | -| `fbuiltins.py:502-520` `as_connectivity_field` | `self.value` (N1) | A1 | -| `nd_array_field.py:981-985` reductions | `axis.value` (N3) | A3 — carries the comment `# assumes offset and local dimension have same name` | -| `nd_array_field.py:972-979` | — | `axis.kind == LOCAL`; at most one local dim per field | -| `nd_array_field.py:317-320` `premap` | — | `FieldOffset` to `Connectivity` via A1 | - -### 4.2 Embedded — iterator level (`iterator/embedded.py`) - -| Site | Key used | Constraint | -| --------------------------------------- | --------------------------------------------------------- | --------------------------------- | -| `:596-616` `execute_shift` | tag (N4), then `source_dim.value` / `codomain.value` (N3) | A7, A8, A9 | -| `:566-576` sparse shift | tag (N4) | A3 | -| `:941-953` `make_in_iterator` | `sparse_dimensions[0].value` (N3) used as tag | A3 | -| `:1517-1519` `SparseListIterator.deref` | `self.list_offset` (N3-derived) | A3 | -| `:1005` `field_setitem` | `value.offset.value` used as a **field dim name** | A3 (tag to N3, reverse direction) | -| `:1410-1416` `_List.__gt_type__` | tag, then `neighbor_dim` | correct direction, no assumption | -| `:1436-1451` `neighbors` | `offset.value` (N1) | A1 | -| `:1776` `_fieldspec_list_to_value` | `offset_type.value` (N3) | A3 | - -### 4.3 IR / type system - -| Site | Key used | Constraint | -| ------------------------------------------------------- | ----------------------------------------------- | ------------------------------------------------------------------------------------ | -| `type_synthesizer.py:326-329` `neighbors` | `OffsetLiteral.value` (N5), then `neighbor_dim` | A2; local dim taken from provider, **not** from the tag | -| `type_synthesizer.py:740-758` `shift` | N5, then `domain[0]`/`codomain` | A2, A7, A8 (`assert found`, `assert not found`) | -| `type_synthesizer.py:433-447` `_canonicalize_nb_fields` | field's LOCAL dim to `ListType.offset_type` | where N3 enters `ListType` and becomes an A3 key downstream | -| `type_synthesizer.py:546-556` `_resolve_dimensions` | N5, then `get_offset_type` | A2 | -| `unroll_reduce.py:43-50, 61-65` | `arg.type.offset_type.value` (N3) | **A3** | -| `domain_utils.py:205-223` | `off.value` (N5) | A2 | -| `pass_manager.py:55-63` | `source_dim.value`/`codomain.value` (N3) | domain sizes keyed by dimension name | -| `past_to_itir.py:409-410` | — | `ValueError: "common.Dimension '{dim.value}' must not be local."` in program domains | -| `type_deduction.py:459-464` | — | `"Second dimension in offset must be a local dimension."` | -| `type_info.py:637-650, 848-878` | — | shift typing via `source`/`target` only; the tag is never consulted | - -### 4.4 GTFN backend - -| Site | Name used | Constraint | -| ------------------------------------- | -------------------------------------------------- | ---------------------------------------------------------- | -| `itir_to_gtfn_ir.py:181-190` | provider key **and** `neighbor_dim.value` | the only site that anticipates A5 failing; emits both tags | -| `itir_to_gtfn_ir.py:191-196` | `source_dim.value`, `codomain.value` | must be `HORIZONTAL`, else `NotImplementedError` | -| `itir_to_gtfn_ir.py:197-200` | — | provider entries must be `NeighborConnectivityType` | -| `itir_to_gtfn_ir.py:485-492` | N5 tags | `o in self.offset_provider_type` | -| `itir_to_gtfn_ir.py:139-148, 166-180` | `dim.value` (N3) | every field dim name becomes a C++ tag | -| `gtfn_module.py:88-98` | `dim.value` (N3) | **A4** | -| `gtfn_module.py:126-136` | `domain[0].value`, `domain[1].value`, provider key | all three become `generated::_t` | - -### 4.5 DaCe backend - -| Site | Name used | Constraint | -| ---------------------------------------------------- | -------------------------------------- | ------------------------------------------------------------------------------------------------------------------- | -| `gtir_to_sdfg.py:572-585` | `local_dim.value` (N3) | **A4**, explicit: `ValueError("The provided local dimension {local_dim} does not match any offset provider type.")` | -| `gtir_to_sdfg.py:838-842` | `dim.value` (N3) | A4 — array shape from `max_neighbors` | -| `gtir_to_sdfg_lambda.py:766-770` | `local_dim.value` (N3) | A4 | -| `gtir_to_sdfg_lambda.py:1312-1319, 1371, 1443, 1455` | `offset_type.value` (N3) | A3, plus connectivity array name | -| `gtir_to_sdfg_lambda.py:1155` | tag (N5) to `Dimension(offset, LOCAL)` | reverse of A5; conflicts with `type_synthesizer.py:329` | -| `gtir_to_sdfg_lambda.py:1718-1727` | `offset_provider_arg.value` (N5) | genuine tag lookup — correct | -| `gtir_to_sdfg_primitives.py:324-331` | `offset_type.value` (N3) | A3 | -| `gtir_to_sdfg_scan.py:385-389` | `offset_type.value` (N3) | A3 | -| `sdfg_args.py:73-93` | field name plus `dim.value` | dim matched against `source_dim`/`neighbor_dim`, else `ValueError` | - -## 5. Constraints on the *format* of names - -| # | Constraint | Source | -| ------ | --------------------------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------------------------- | -| **F1** | `_Staggered` is a **reserved prefix**: any `Dimension` whose `value` starts with it is treated as staggered | `common.py:1444-1464` (`_STAGGERED_PREFIX = "_Staggered"`) | -| **F2** | GTFN aliases every staggered tag to its base tag by string surgery | `itir_to_gtfn_ir.py:703`, `_add_staggered_aliases:204-215` | -| **F3** | `_CONST_DIM` is a reserved LOCAL dimension name, deliberately *absent* from the provider and special-cased at every lookup | `embedded.py:220, 572, 1513, 1768`; `gtir_to_sdfg_lambda.py:62, 1314, 1355` | -| **F4** | **Dimension names determine memory layout** — `order_dimensions` sorts by `(kind, as_non_staggered(dim).value)` | `common.py:1334-1344` | -| **F5** | GTFN: every dim name and provider key becomes a C++ type `generated::_t`, so it must be a valid C++ identifier and collision-free | `gtfn_module.py:97, 130-136` | -| **F6** | GTFN connectivity params: `gt_conn_`, so keys must not collide **case-insensitively** | `gtfn_module.py:32, 120, 133` | -| **F7** | DaCe connectivity arrays: `gt_conn_`, recovered by regex `^gt_conn_(\S+)$` | `sdfg_args.py:24-25, 56-70` | -| **F8** | DaCe map variables: `i__gtx_[dim]`; map fusion/splitting transformations **rely on these strings matching** | `gtir_to_sdfg_utils.py:44-54` | -| **F9** | DaCe field symbols: `____size/stride`, `_range_symbol_name(field, dim.value)` | `sdfg_args.py:73-82, 119-122` | - -## 6. Structural (kind / arity) constraints - -| # | Constraint | Enforced | Source | -| --- | ------------------------------------------------------------------------- | ------------------------------------ | ----------------------------------------------------------------------------- | -| S1 | `len(target) == 2` implies `target[1].kind == LOCAL` | eager `ValueError` | `fbuiltins.py:480-482`; also `type_deduction.py:459-462` | -| S2 | A neighbor table's domain is exactly `(HORIZONTAL, LOCAL)` | `is_neighbor_table` guard | `common.py:1160-1168` | -| S3 | At most one LOCAL dim per field | `ValueError` / `NotImplementedError` | `common.py:1334-1337`; `nd_array_field.py:976-979`; `gtir_to_sdfg.py:586-589` | -| S4 | Cartesian offset iff `len(target)==1 and source==target[0]` and not LOCAL | predicate | `fbuiltins.py:524-529` | -| S5 | Non-Cartesian offset or LOCAL dim implies grid type `UNSTRUCTURED` | `ValueError` | `transform_utils.py:60-77` | -| S6 | `as_offset` is Cartesian-only | `DSLError` | `type_deduction.py:955-965` | -| S7 | Program domains must not contain LOCAL dims | `ValueError` | `past_to_itir.py:409-410` | - -## 7. Observed behaviour - -Two properties above were confirmed by running them, not only by reading. - -### 7.1 A1 vs A2 — embedded and compiled key on different strings - -```python -MyOff = gtx.FieldOffset("TAGNAME", source=E, target=(V, Neigh)) # tag != variable name - - -@gtx.field_operator -def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return a(MyOff[1]) -``` - -```text -embedded: offset_provider={"TAGNAME": conn} -> OK ; {"MyOff": conn} -> KeyError 'TAGNAME' -roundtrip: offset_provider={"MyOff": conn} -> OK ; {"TAGNAME": conn} -> KeyError 'MyOff' -``` - -The compiled path uses the Python variable name because `foast_to_gtir.py:302-306` -and `:325-331` emit `im.shift(offset_name.id, ...)` / `im.as_fieldop_neighbors(str(offset_name), ...)` -from the FOAST `Name.id` — never from `FieldOffset.value`. Lowering the operator -above yields: - -```text -foo = λ(a) → (⇑(λ(__it) → ·⟪MyOffₒ, 1ₒ⟫(__it)))(a); -``` - -### 7.2 A3 — reductions still require tag == local dim name - -Reusing the deliberately mismatched declaration from the #1789 regression test -(`Off` tagged `"Off"`, local dim named `"Neigh"`): - -```python -Off = gtx.FieldOffset("Off", source=E, target=(V, Neigh)) - - -@gtx.field_operator -def bar(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return neighbor_sum(a(Off), axis=Neigh) -``` - -```text -embedded: FAILED: KeyError: "Offset 'Neigh' not found in offset provider." -roundtrip: OK -> [30. 50. 40.] -``` - -`unroll_reduce.py:43-50` has the same assumption for the compiled pipeline -(established by reading; the roundtrip backend above does not exercise that pass). - -## 8. Practical consequence - -To be safe across **all** contexts, four strings must be identical: - -```text -FieldOffset.value == == offset_provider key == target[-1].value -``` - -plus `target[0] == conn.domain[0]` and `source == conn.codomain` as `Dimension` -objects (A6-A8). This is exactly what `tests/next_tests/toy_connectivity.py:18-26` -encodes: - -```python -V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) # value is "V2E", not "V2EDim" -V2E = gtx.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) -``` - -Relaxing any one of the four is currently supported only in the narrow slice -PR #1789 covered: shift-only, GTFN, no sparse arguments. Nothing validates the -full set up front — a violation surfaces as a `KeyError` from `common.py:1208`, -a bare `assert`, or, per the #1789 test docstring, silently wrong results. - -Two existing `TODO`s point at this tangle: - -- `common.py:976-977` — `NeighborConnectivityType`: *"refactor towards encoding - this information in the local dimensions of the `ConnectivityType.domain`"*. -- `fbuiltins.py:467-470` — *"`FieldOffset` and `runtime.Offset` are not an exact - conceptual match. Revisit if we want to continue subclassing here."* diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index b479a81326..f65b7bc0b6 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -373,7 +373,11 @@ def resolve(tag: Tag) -> Dimension: """ if (match := _STAGGERED_TAG_RE.match(tag)) is not None: owner = resolve(match["owner"]) - return owner[resolve(match["base"])] # type: ignore[index] # parametrized dimension + if not isinstance(owner, StaggeredMeta): + raise ValueError( + f"Cannot resolve tag '{tag}': '{match['owner']}' is not a parametrized dimension." + ) + return owner[resolve(match["base"])] # type: ignore[index] # a StaggeredMeta, checked parts = tag.split(".") for split in range(len(parts), 0, -1): @@ -1725,8 +1729,12 @@ def __getitem__(cls, base: Dimension) -> Dimension: ) if not isinstance(base, DimensionMeta): raise TypeError(f"'Staggered' expects a dimension, got '{base!r}'.") - if base not in _STAGGERED_CACHE: - _STAGGERED_CACHE[base] = cast( + if is_staggered(base): + raise TypeError( + f"'{base.__qualname__}' is already staggered; a dimension cannot be staggered twice." + ) + if (staggered := _STAGGERED_CACHE.get(base)) is None: + staggered = cast( Dimension, StaggeredMeta( f"Staggered[{base.__qualname__}]", @@ -1744,7 +1752,10 @@ def __getitem__(cls, base: Dimension) -> Dimension: }, ), ) - return _STAGGERED_CACHE[base] + # NOTE: `setdefault`, not an assignment: compilation runs in threads, and two of them + # building `Staggered[K]` at once must still see one class (identity is the dimension). + staggered = _STAGGERED_CACHE.setdefault(base, staggered) + return staggered if TYPE_CHECKING: diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index b439b7d77e..3d09d7ecb4 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -482,7 +482,9 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri " choose one." ) ], - hints=[f"Write the displacement directly, e.g. '{source.value} + 1'."], + hints=[ + f"Write the displacement directly, e.g. '{source.__qualname__} + 1'." + ], ) new_type = new_value.type case ts.FieldType(dims=dims, dtype=dtype): diff --git a/src/gt4py/next/otf/runners.py b/src/gt4py/next/otf/runners.py index 75b54f19a3..8ee324c6d5 100644 --- a/src/gt4py/next/otf/runners.py +++ b/src/gt4py/next/otf/runners.py @@ -208,7 +208,9 @@ def reducer_override(self, o: object) -> Any: try: _Scanner(io.BytesIO()).dump(obj) - except Exception: # an unpicklable object is reported by the executor check instead + except Exception: + # An object that cannot be pickled at all: not this check's business. It surfaces when + # the pool pickles the task, which is where an unpicklable job is reported. return None return found[0] if found else None diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index dcdc6128d2..500a355dbb 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -900,3 +900,20 @@ def test_gt_dims_are_unqualified_names(): """ field = gtx.as_field([IDim, JDim], np.zeros((2, 3))) assert field.__gt_dims__ == ("IDim", "JDim") + + +class TestStaggered: + def test_interned(self): + assert common.Staggered[KDim] is common.Staggered[KDim] + assert common.is_staggered(common.Staggered[KDim]) + assert common.flip_staggered(common.Staggered[KDim]) is KDim + + def test_a_dimension_cannot_be_staggered_twice(self): + with pytest.raises(TypeError, match="is already staggered"): + common.Staggered[common.Staggered[KDim]] + with pytest.raises(TypeError, match="is already staggered"): + common.Staggered[KDim][KDim] + + def test_resolve_rejects_a_bracketed_tag_of_another_owner(self): + with pytest.raises(ValueError, match="not a parametrized dimension"): + common.resolve(f"{KDim.tag}[{KDim.tag}]") From 996a34ab27196084fa747b97d05e457faa17d01a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 04:54:51 +0200 Subject: [PATCH 13/20] refactor[next]: ConstList replaces _CONST_DIM; AxisLiteral drops kind MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - common.ConstListDim becomes common.ConstList, the local dimension of make_const_list, used directly by iterator embedded and the DaCe lowering (identity checks) - AxisLiteral stores only the tag; kind and dim are resolved from it. The pretty printer derives the suffix (now with ₗ for local dimensions, which used to be printed as vertical), and the parser ignores it. --- .../next/0028-Dimensions_As_Nominal_Types.md | 4 ++- src/gt4py/next/common.py | 16 +++++------ src/gt4py/next/ffront/foast_to_gtir.py | 4 +-- src/gt4py/next/ffront/past_to_itir.py | 6 ++--- src/gt4py/next/iterator/embedded.py | 18 +++++-------- src/gt4py/next/iterator/ir.py | 15 ++++++++--- src/gt4py/next/iterator/ir_utils/ir_makers.py | 4 +-- src/gt4py/next/iterator/pretty_parser.py | 7 +++-- src/gt4py/next/iterator/pretty_printer.py | 16 ++++++++--- src/gt4py/next/iterator/tracing.py | 2 +- .../iterator/transforms/remove_broadcast.py | 4 +-- .../dace/lowering/gtir_to_sdfg_lambda.py | 24 +++++++---------- .../ffront_tests/test_foast_to_gtir.py | 4 +-- .../test_embedded_field_with_list.py | 5 ++-- .../iterator_tests/test_pretty_parser.py | 4 +-- .../iterator_tests/test_pretty_printer.py | 27 ++++++++++++------- .../iterator_tests/test_pretty_roundtrip.py | 17 +++--------- .../iterator_tests/test_type_inference.py | 20 +++++--------- 18 files changed, 93 insertions(+), 104 deletions(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 224f8ec289..295ceebe59 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -78,7 +78,9 @@ disappears. 1. **Reconstruction from the IR is an import.** `common.resolve(tag)` imports the module and walks the qualname; nested declarations resolve naturally. The IR references a Python type exactly the way `pickle` references a class. It is - memoized, because type inference calls it once per `AxisLiteral`. + memoized, because type inference calls it once per `AxisLiteral`. An + `AxisLiteral` stores only the tag: its `kind` is the resolved dimension's, so the + two cannot disagree. A purely dotted tag does not record where the module path ends and the qualname begins, so `resolve` tries the *longest importable prefix* and walks diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index f65b7bc0b6..a9ced01a37 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1789,19 +1789,15 @@ def __init_subclass__(cls, /, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) -class ConstListDim(DimensionIndex, kind=DimensionKind.LOCAL): +class ConstList(DimensionIndex, kind=DimensionKind.LOCAL): """ - The local dimension of a list whose length is known at compile time (`make_const_list`). + The local dimension of a list of one repeated value (`make_const_list`). - Declared here, once, because it must be a *single* class. It used to be built - independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless - while dimensions compared by `(name, kind)` -- the two instances were equal. Under - nominal identity (ADR 0028) two declarations would be two different dimensions, and the - `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s - built by embedded execution. + The value is broadcast against the neighbor lists it is combined with, and a materialized + constant list has extent 1 along it. It indexes no table, so it is never in an offset provider. - TODO: becomes an owner-less local dimension with an explicit size, generalising this from - length 1 to length *n*, once local dimensions know their connectivity. + Declared here, once: it used to be built independently in `iterator/embedded.py` and in the + DaCe lowering, which only worked while dimensions compared by `(name, kind)`. """ __slots__ = () diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 4817072a26..929650cf02 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -233,12 +233,12 @@ def visit_Symbol(self, node: foast.Symbol, **kwargs: Any) -> itir.Sym: def visit_Name(self, node: foast.Name, **kwargs: Any) -> itir.SymRef | itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag) return im.ref(node.id) def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag) if isinstance(named_tup_type := node.value.type, ts.NamedCollectionType): ind = named_tup_type.keys.index(node.attr) diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 1b6cefc6b2..bdedbc3da9 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -383,9 +383,7 @@ def _construct_itir_domain_arg( domain_args = [] for dim_i, dim in enumerate(out_type.dims): # an expression for the range of a dimension - dim_range = im.call("get_domain_range")( - out_expr, itir.AxisLiteral(value=dim.tag, kind=dim.kind) - ) + dim_range = im.call("get_domain_range")(out_expr, itir.AxisLiteral(value=dim.tag)) dim_start, dim_stop = im.tuple_get(0, dim_range), im.tuple_get(1, dim_range) # bounds @@ -413,7 +411,7 @@ def _construct_itir_domain_arg( domain_args.append( itir.FunCall( fun=itir.SymRef(id="named_range"), - args=[itir.AxisLiteral(value=dim.tag, kind=dim.kind), lower, upper], + args=[itir.AxisLiteral(value=dim.tag), lower, upper], ) ) diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 32f31c29a9..895215d0d7 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -211,12 +211,6 @@ def skip_value( NamedFieldIndices: TypeAlias = Mapping[Tag, FieldIndex | SparsePositionEntry] -# Magic local dimension for the result of a `make_const_list`. -# A clean implementation will probably involve to tag the `make_const_list` -# with the neighborhood it is meant to be used with. -_CONST_DIM = common.ConstListDim - - @runtime_checkable class ItIterator(Protocol): """ @@ -565,7 +559,7 @@ def execute_shift( for i, p in reversed(list(enumerate(new_entry))): # first shift applies to the last sparse dimensions of that axis type if p is None: - if tag == _CONST_DIM.tag: + if tag == common.ConstList.tag: new_entry[i] = 0 else: offset_implementation = common.get_offset(offset_provider, tag) @@ -1003,7 +997,7 @@ def field_setitem(self, named_indices: NamedFieldIndices, value: Any): ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ - self._translate_named_indices({**named_indices, _CONST_DIM.tag: 0}) + self._translate_named_indices({**named_indices, common.ConstList.tag: 0}) ] = value.value else: self._ndarrayfield[self._translate_named_indices(named_indices)] = value @@ -1430,7 +1424,7 @@ def __gt_type__(self) -> ts.ListType: assert isinstance(element_type, ts.DataType) return ts.ListType( element_type=element_type, - offset_type=_CONST_DIM, + offset_type=common.ConstList, ) @@ -1512,7 +1506,7 @@ class SparseListIterator: offsets: Sequence[OffsetPart] = dataclasses.field(default_factory=list, kw_only=True) def deref(self) -> Any: - if self.list_offset == _CONST_DIM.tag: + if self.list_offset == common.ConstList.tag: return _ConstList( value=self.it.shift(*self.offsets, SparseTag(self.list_offset), 0).deref() ) @@ -1767,9 +1761,9 @@ def _fieldspec_list_to_value( ) -> tuple[common.Domain, ts.TypeSpec]: """Translate the list element type into the domain.""" if isinstance(type_, ts.ListType): - if type_.offset_type == _CONST_DIM: + if type_.offset_type is common.ConstList: return domain.insert( - len(domain), common.named_range((_CONST_DIM, 1)) + len(domain), common.named_range((common.ConstList, 1)) ), type_.element_type else: offset_provider = embedded_context.get_offset_provider() diff --git a/src/gt4py/next/iterator/ir.py b/src/gt4py/next/iterator/ir.py index f024ef6168..05561c2398 100644 --- a/src/gt4py/next/iterator/ir.py +++ b/src/gt4py/next/iterator/ir.py @@ -90,10 +90,19 @@ class OffsetLiteral(Expr): class AxisLiteral(Expr): - # TODO(havogt): Refactor to use declare Axis/Dimension at the Program level. - # Now every use of the literal has to provide the kind, where usually we only care of the name. + #: The dimension's tag, its qualified Python name (ADR 0028). value: str - kind: common.DimensionKind = common.DimensionKind.HORIZONTAL + + @property + def dim(self) -> common.Dimension: + """The dimension the literal names, resolved from its tag.""" + return common.resolve(self.value) + + @property + def kind(self) -> common.DimensionKind: + # NOTE: derived, not stored: the dimension class carries its kind, so a stored copy could + # only disagree with it (it used to, for local dimensions printed as vertical). + return self.dim.kind class CartesianOffset(Expr): diff --git a/src/gt4py/next/iterator/ir_utils/ir_makers.py b/src/gt4py/next/iterator/ir_utils/ir_makers.py index d6458f7a88..053e7b78c8 100644 --- a/src/gt4py/next/iterator/ir_utils/ir_makers.py +++ b/src/gt4py/next/iterator/ir_utils/ir_makers.py @@ -583,7 +583,7 @@ def _impl(*its: itir.Expr) -> itir.FunCall: def axis_literal(dim: common.Dimension) -> itir.AxisLiteral: - return itir.AxisLiteral(value=dim.tag, kind=dim.kind) + return itir.AxisLiteral(value=dim.tag) def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: @@ -640,7 +640,7 @@ def index(dim: common.Dimension) -> itir.FunCall: Returns: A function that constructs a Field of indices in the given dimension. """ - return call("index")(itir.AxisLiteral(value=dim.tag, kind=dim.kind)) + return call("index")(itir.AxisLiteral(value=dim.tag)) def map_list(op): diff --git a/src/gt4py/next/iterator/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index 07033ff1db..1fd5f557bc 100644 --- a/src/gt4py/next/iterator/pretty_parser.py +++ b/src/gt4py/next/iterator/pretty_parser.py @@ -41,7 +41,7 @@ // suffix terminates it. TAG: /[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*(?:\[[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*\])?/ OFFSET_LITERAL: ( INT_LITERAL | TAG ) "ₒ" - AXIS_LITERAL: TAG ("ᵥ" | "ₕ") + AXIS_LITERAL: TAG ("ᵥ" | "ₕ" | "ₗ") INFINITY_LITERAL: "∞" | "-∞" _literal: INT_LITERAL | FLOAT_LITERAL | OFFSET_LITERAL | AXIS_LITERAL | INFINITY_LITERAL ID_NAME: CNAME @@ -177,9 +177,8 @@ def INFINITY_LITERAL(self, value: lark_lexer.Token) -> ir.InfinityLiteral: return ir.InfinityLiteral.POSITIVE def AXIS_LITERAL(self, value: lark_lexer.Token) -> ir.AxisLiteral: - name = value.value[:-1] - kind = ir.DimensionKind.HORIZONTAL if value.value[-1] == "ₕ" else ir.DimensionKind.VERTICAL - return ir.AxisLiteral(value=name, kind=kind) + # NOTE: the kind suffix is only for the reader; the kind is the dimension's own. + return ir.AxisLiteral(value=value.value[:-1]) def lam(self, *args: ir.Node) -> ir.Lambda: *params, expr = args diff --git a/src/gt4py/next/iterator/pretty_printer.py b/src/gt4py/next/iterator/pretty_printer.py index 5fbba8920e..5321c01089 100644 --- a/src/gt4py/next/iterator/pretty_printer.py +++ b/src/gt4py/next/iterator/pretty_printer.py @@ -19,6 +19,7 @@ from typing import Final from gt4py.eve import NodeTranslator +from gt4py.next import common from gt4py.next.iterator import ir from gt4py.next.type_system import type_specifications as ts, type_translation @@ -134,6 +135,13 @@ def implied_literal_type(value: str) -> ts.ScalarType: DEFAULT_WIDTH: Final = 100 +_AXIS_KIND_SUFFIX: Final = { + common.DimensionKind.HORIZONTAL: "ₕ", + common.DimensionKind.VERTICAL: "ᵥ", + common.DimensionKind.LOCAL: "ₗ", +} + + class PrettyPrinter(NodeTranslator): def __init__( self, @@ -225,11 +233,11 @@ def visit_CartesianOffset(self, node: ir.CartesianOffset, *, prec: int) -> list[ return [f"{domain}→{codomain}"] def visit_AxisLiteral(self, node: ir.AxisLiteral, *, prec: int) -> list[str]: - kind = "" - if node.kind == ir.DimensionKind.HORIZONTAL: + try: + kind = _AXIS_KIND_SUFFIX[node.kind] + except ValueError: + # a tag that names no importable dimension, e.g. in IR built by hand for debugging kind = "ₕ" - elif node.kind == ir.DimensionKind.VERTICAL: - kind = "ᵥ" return [str(node.value) + kind] def visit_SymRef(self, node: ir.SymRef, *, prec: int) -> list[str]: diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 4e5b8d4b5a..a6ae605bab 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -141,7 +141,7 @@ def make_node(o): if isinstance(o, Node): return o if isinstance(o, common.DimensionMeta): - return AxisLiteral(value=o.tag, kind=o.kind) + return AxisLiteral(value=o.tag) if isinstance(o, common.Infinity): if o is common.Infinity.POSITIVE: return itir.InfinityLiteral.POSITIVE diff --git a/src/gt4py/next/iterator/transforms/remove_broadcast.py b/src/gt4py/next/iterator/transforms/remove_broadcast.py index db6445cc51..1a92361216 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -34,9 +34,7 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): >>> domain = im.domain(common.GridType.CARTESIAN, {IDim: (0, 10), JDim: (0, 10)}) >>> expr = im.call("broadcast")( ... im.ref("inp"), - ... im.make_tuple( - ... *(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (IDim, JDim)) - ... ), + ... im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (IDim, JDim))), ... ) >>> expr.annex.domain = domain_utils.SymbolicDomain.from_expr(domain) >>> transformed = RemoveBroadcast.apply(expr) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index ed3ff892ec..b5a3a1d632 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -58,13 +58,6 @@ ) -# Magic local dimension used for list of values with length known at compile-time. -# NOTE: the canonical class from `common`, not a local declaration: under nominal identity a -# second declaration would be a *different* dimension and the `== _CONST_DIM` checks below -# would stop matching `ListType`s built by embedded execution. -_CONST_DIM: Final = gtx_common.ConstListDim - - @dataclasses.dataclass(frozen=True) class ValueExpr: """ @@ -595,7 +588,7 @@ def _construct_tasklet_result( return ValueExpr( dc_node=temp_node, gt_dtype=( - ts.ListType(element_type=data_type, offset_type=_CONST_DIM) + ts.ListType(element_type=data_type, offset_type=gtx_common.ConstList) if use_array else data_type ), @@ -1142,7 +1135,8 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: gt_field=ts.FieldType( dims=[conn_type.domain[0]], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1318,7 +1312,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: assert isinstance(input_arg.gt_dtype, ts.ListType) assert input_arg.gt_dtype.offset_type is not None offset_type = input_arg.gt_dtype.offset_type - if offset_type == _CONST_DIM: + if offset_type is gtx_common.ConstList: # this input argument is the result of `make_const_list` continue offset_provider_t = self.subgraph_builder.get_offset_provider_type(offset_type.tag) @@ -1359,7 +1353,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: raise ValueError(f"More than one local dimension in map expression {node}.") input_size = input_desc.shape[0] if input_size == 1: - assert input_arg.gt_dtype.offset_type == _CONST_DIM + assert input_arg.gt_dtype.offset_type is gtx_common.ConstList input_memlets[conn] = dace.Memlet(data=input_node.data, subset="0") elif input_size == local_size: input_memlets[conn] = dace.Memlet(data=input_node.data, subset=map_index) @@ -1390,7 +1384,8 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: gt_field=ts.FieldType( dims=[conn_type.domain[0]], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1716,7 +1711,8 @@ def _make_unstructured_shift( gt_field=ts.FieldType( dims=[conn_type.source_dim], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1921,7 +1917,7 @@ def _visit_Lambda_impl( and node.expr.type.offset_type is not None and isinstance(result, (MemletExpr, ValueExpr)) and isinstance(result.gt_dtype, ts.ListType) - and result.gt_dtype.offset_type == _CONST_DIM + and result.gt_dtype.offset_type is gtx_common.ConstList ): result = self._broadcast_const_list(result, node.expr.type) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index f696daf1b4..0e2cf62732 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 @@ -985,7 +985,7 @@ def foo(inp: gtx.Field[[TDim], float64]): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( im.ref("inp"), - im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (TDim, UDim))), ) @@ -999,7 +999,7 @@ def foo(): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( 1, - im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (TDim, UDim))), ) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py index 12c0b75649..869372ab78 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py @@ -10,6 +10,7 @@ import pytest import gt4py.next as gtx +from gt4py.next import common from gt4py.next.embedded import context as embedded_context from gt4py.next.iterator import embedded, runtime from gt4py.next.iterator.builtins import ( @@ -68,7 +69,7 @@ def testee(): ref = np.asarray([[42.0], [42.0]]) assert result.domain.dims[0] == E - assert result.domain.dims[1] == embedded._CONST_DIM # this is implementation detail + assert result.domain.dims[1] == common.ConstList # this is implementation detail assert result.shape[1] == 1 # this is implementation detail np.testing.assert_array_equal(result.asnumpy(), ref) @@ -151,6 +152,6 @@ def testee(): ref = np.asarray([[43.0], [43.0]]) assert result.domain.dims[0] == E - assert result.domain.dims[1] == embedded._CONST_DIM # this is implementation detail + assert result.domain.dims[1] == common.ConstList # this is implementation detail assert result.shape[1] == 1 # this is implementation detail np.testing.assert_array_equal(result.asnumpy(), ref) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py index 83a16a980c..52b358023d 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py @@ -190,7 +190,7 @@ def test_named_range_unbounded(): expected = ir.FunCall( fun=ir.SymRef(id="named_range"), args=[ - ir.AxisLiteral(value="KDim", kind=ir.DimensionKind.VERTICAL), + ir.AxisLiteral(value="KDim"), ir.InfinityLiteral.NEGATIVE, ir.InfinityLiteral.POSITIVE, ], @@ -257,7 +257,7 @@ def test_named_range_vertical(): expected = ir.FunCall( fun=ir.SymRef(id="named_range"), args=[ - ir.AxisLiteral(value="IDim", kind=ir.DimensionKind.VERTICAL), + ir.AxisLiteral(value="IDim"), ir.SymRef(id="x"), ir.SymRef(id="y"), ], diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py index 1144b727b2..97d1805f69 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py @@ -8,6 +8,8 @@ import pytest +import gt4py.next as gtx + from gt4py.next.iterator import builtins, ir, pretty_printer from gt4py.next.iterator.ir_utils import ir_makers as im from gt4py.next.iterator.pretty_printer import PrettyPrinter, pformat @@ -259,18 +261,23 @@ def test_make_tuple(): assert actual == expected -def test_axis_literal_horizontal(): - testee = ir.AxisLiteral(value="I", kind=ir.DimensionKind.HORIZONTAL) - expected = "Iₕ" - actual = pformat(testee) - assert actual == expected +class IDim(gtx.DimensionIndex): ... -def test_axis_literal_vertical(): - testee = ir.AxisLiteral(value="I", kind=ir.DimensionKind.VERTICAL) - expected = "Iᵥ" - actual = pformat(testee) - assert actual == expected +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + + +@pytest.mark.parametrize("dim, suffix", [(IDim, "ₕ"), (KDim, "ᵥ"), (LocalDim, "ₗ")]) +def test_axis_literal(dim, suffix): + # the suffix is the resolved dimension's kind + assert pformat(ir.AxisLiteral(value=dim.tag)) == f"{dim.tag}{suffix}" + + +def test_axis_literal_of_unresolvable_tag(): + assert pformat(ir.AxisLiteral(value="I")) == "Iₕ" def test_named_range_horizontal(): diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py index 81396d8c59..6584028a93 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_roundtrip.py @@ -72,19 +72,8 @@ im.tuple_get(im.literal("42", builtins.INTEGER_INDEX_BUILTIN), "x"), id="tuple_get" ), pytest.param(im.make_tuple("x", "y"), id="make_tuple"), - pytest.param( - ir.AxisLiteral(value="I", kind=ir.DimensionKind.HORIZONTAL), id="axis_literal_horizontal" - ), - pytest.param( - ir.AxisLiteral(value="I", kind=ir.DimensionKind.VERTICAL), id="axis_literal_vertical" - ), - pytest.param( - im.named_range(ir.AxisLiteral(value="IDim"), "x", "y"), id="named_range_horizontal" - ), - pytest.param( - im.named_range(ir.AxisLiteral(value="IDim", kind=ir.DimensionKind.VERTICAL), "x", "y"), - id="named_range_vertical", - ), + pytest.param(ir.AxisLiteral(value="I"), id="axis_literal"), + pytest.param(im.named_range(ir.AxisLiteral(value="IDim"), "x", "y"), id="named_range"), pytest.param(im.call("cartesian_domain")("x"), id="cartesian_domain"), pytest.param(im.call("unstructured_domain")("x"), id="unstructured_domain"), pytest.param(im.if_("x", "y", "z"), id="if_short"), @@ -161,7 +150,7 @@ pytest.param(ir.InfinityLiteral.NEGATIVE, id="infinity_negative"), pytest.param( im.named_range( - ir.AxisLiteral(value="KDim", kind=ir.DimensionKind.VERTICAL), + ir.AxisLiteral(value="KDim"), ir.InfinityLiteral.NEGATIVE, im.literal("5", "int32"), ), diff --git a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py index 94f0610c06..4f580b110c 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py @@ -101,9 +101,7 @@ def expression_test_cases(): bool_type, ), ( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), it_ts.NamedRangeType(dim=Vertex), ), ( @@ -112,9 +110,7 @@ def expression_test_cases(): ), ( im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ) + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1) ), ts.DomainType(dims=[Vertex]), ), @@ -443,10 +439,8 @@ def test_cartesian_fencil_definition(): def test_unstructured_fencil_definition(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), - im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( @@ -510,10 +504,8 @@ def test_function_definition(): def test_fencil_with_nb_field_input(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), - im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( From 83c362f92b8dc591e0fe4d9bc3b932dce677f944 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 06:28:30 +0200 Subject: [PATCH 14/20] fix[next]: printing IR does not import modules The pretty printer takes an axis literal's kind from its inferred type, or from a dimension that is already loaded (common.resolve_loaded), and never imports. gtfn's domain canonicalization resolves through dim_from_axis_literal. --- src/gt4py/next/common.py | 20 +++++++++++++++++++ src/gt4py/next/iterator/pretty_parser.py | 2 +- src/gt4py/next/iterator/pretty_printer.py | 15 ++++++++------ .../codegens/gtfn/itir_to_gtfn_ir.py | 9 +++++++-- .../iterator_tests/test_pretty_parser.py | 3 ++- .../iterator_tests/test_pretty_printer.py | 8 ++++++++ 6 files changed, 47 insertions(+), 10 deletions(-) diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index a9ced01a37..5839d2bbe0 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -402,6 +402,26 @@ def resolve(tag: Tag) -> Dimension: ) +def resolve_loaded(tag: Tag) -> Optional[Dimension]: + """ + Return the dimension a tag names if its module is already loaded, else `None`. + + Like `resolve`, but never imports: for code that must not have import side effects, such as + printing IR. + """ + if (match := _STAGGERED_TAG_RE.match(tag)) is not None: + owner, base = resolve_loaded(match["owner"]), resolve_loaded(match["base"]) + return owner[base] if owner is not None and base is not None else None # type: ignore[index] # parametrized dimension + parts = tag.split(".") + for split in range(len(parts) - 1, 0, -1): + if (obj := sys.modules.get(".".join(parts[:split]))) is None: + continue + for attr in parts[split:]: + obj = getattr(obj, attr, None) + return obj if isinstance(obj, DimensionMeta) else None + return None + + class Infinity(enum.Enum): """Describes an unbounded `UnitRange`.""" diff --git a/src/gt4py/next/iterator/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index 1fd5f557bc..0324499d26 100644 --- a/src/gt4py/next/iterator/pretty_parser.py +++ b/src/gt4py/next/iterator/pretty_parser.py @@ -20,7 +20,7 @@ from gt4py.next.type_system import type_specifications as ts -GRAMMAR = """ +GRAMMAR = r""" start: fencil_definition | function_definition | declaration diff --git a/src/gt4py/next/iterator/pretty_printer.py b/src/gt4py/next/iterator/pretty_printer.py index 5321c01089..a7376e0d22 100644 --- a/src/gt4py/next/iterator/pretty_printer.py +++ b/src/gt4py/next/iterator/pretty_printer.py @@ -16,7 +16,7 @@ import types as _types from collections.abc import Iterator, Mapping, Sequence -from typing import Final +from typing import Final, Optional from gt4py.eve import NodeTranslator from gt4py.next import common @@ -233,11 +233,14 @@ def visit_CartesianOffset(self, node: ir.CartesianOffset, *, prec: int) -> list[ return [f"{domain}→{codomain}"] def visit_AxisLiteral(self, node: ir.AxisLiteral, *, prec: int) -> list[str]: - try: - kind = _AXIS_KIND_SUFFIX[node.kind] - except ValueError: - # a tag that names no importable dimension, e.g. in IR built by hand for debugging - kind = "ₕ" + # NOTE: printing must not import modules (`str()` of any node prints it), so the kind is + # taken from the inferred type, or from an already loaded dimension. A tag naming neither, + # e.g. in IR built by hand, prints as horizontal; the parser ignores the suffix anyway. + if isinstance(node.type, ts.DimensionType): + dim: Optional[common.Dimension] = node.type.dim + else: + dim = common.resolve_loaded(node.value) + kind = _AXIS_KIND_SUFFIX[dim.kind] if dim is not None else "ₕ" return [str(node.value) + kind] def visit_SymRef(self, node: ir.SymRef, *, prec: int) -> list[str]: diff --git a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py index 74684d420b..38fe8fa1c6 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py @@ -254,11 +254,16 @@ def visit_FunCall(self, node: itir.FunCall) -> itir.FunCall: assert isinstance(node.args[0], itir.FunCall) first_axis_literal = node.args[0].args[0] assert isinstance(first_axis_literal, itir.AxisLiteral) - if first_axis_literal.kind == itir.DimensionKind.VERTICAL: + if ir_utils_misc.dim_from_axis_literal(first_axis_literal).kind == ( + itir.DimensionKind.VERTICAL + ): assert len(node.args) == 2 assert isinstance(node.args[1], itir.FunCall) assert isinstance(node.args[1].args[0], itir.AxisLiteral) - assert node.args[1].args[0].kind == itir.DimensionKind.HORIZONTAL + assert ( + ir_utils_misc.dim_from_axis_literal(node.args[1].args[0]).kind + == itir.DimensionKind.HORIZONTAL + ) return itir.FunCall(fun=node.fun, args=[node.args[1], node.args[0]]) return node diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py index 52b358023d..7165a54ec2 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py @@ -252,7 +252,8 @@ def test_named_range_horizontal(): assert actual == expected -def test_named_range_vertical(): +def test_named_range_kind_suffix_is_ignored(): + # the kind is the dimension's own; the suffix only helps the reader testee = "IDimᵥ: [x, y[" expected = ir.FunCall( fun=ir.SymRef(id="named_range"), diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py index 97d1805f69..923bda3ff4 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py @@ -277,7 +277,15 @@ def test_axis_literal(dim, suffix): def test_axis_literal_of_unresolvable_tag(): + # printing does not import modules: a tag naming no loaded dimension prints as horizontal, + # so text -> IR -> text is not the identity for such tags (the parser ignores the suffix) assert pformat(ir.AxisLiteral(value="I")) == "Iₕ" + assert pformat(ir.AxisLiteral(value="this.I")) == "this.Iₕ" + + +def test_axis_literal_kind_from_type(): + typed = ir.AxisLiteral(value="not.loaded.KDim", type=ts.DimensionType(dim=KDim)) + assert pformat(typed) == "not.loaded.KDimᵥ" def test_named_range_horizontal(): From 3371a943f24cc6bc5cff9d5480c489af9ffd367d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 07:19:43 +0200 Subject: [PATCH 15/20] test[next]: register doctest dimensions where their tags point --- .../transforms/replace_get_domain_range_with_constants.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py index 26c54abf94..e7626bac47 100644 --- a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py +++ b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py @@ -56,6 +56,9 @@ class ReplaceGetDomainRangeWithConstants(PreserveLocationVisitor, NodeTranslator >>> from gt4py import next as gtx >>> class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... >>> class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + >>> import sys # register the dimensions where their tags point, as a module would + >>> sys.modules[__name__].KDim = KDim + >>> sys.modules[__name__].Vertex = Vertex >>> sizes = { ... "out": gtx.domain({Vertex: (0, 10), KDim: (0, 20)}), From dc53f1a1d180ea780a451791674df76c1596da56 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 16:14:44 +0200 Subject: [PATCH 16/20] fix[next]: review fixes for ConstList and AxisLiteral - ADR 0028's date, which this PR's edits had outrun - say why 'ListType.offset_type' can be 'None' while embedded uses 'ConstList' - test 'resolve_loaded', which is what keeps printing IR import-free --- .../ADRs/next/0028-Dimensions_As_Nominal_Types.md | 2 +- src/gt4py/next/type_system/type_specifications.py | 4 ++++ tests/next_tests/unit_tests/test_common.py | 9 +++++++++ 3 files changed, 14 insertions(+), 1 deletion(-) diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 295ceebe59..d3611694f4 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -7,7 +7,7 @@ tags: [] - **Status**: proposed - **Authors**: Enrique González Paredes (@egparedes) - **Created**: 2026-09-18 -- **Updated**: 2026-09-18 +- **Updated**: 2026-09-24 A concrete dimension becomes a **class**, and an index along it an **instance** of that class — the shape `enum.Enum` uses, where the class is the collection and diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 806350542e..5ad8ea7105 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -114,6 +114,10 @@ class ListType(DataType): """ element_type: DataType + #: The local dimension the list runs along. `None` where type inference does not know it, + #: which is how it spells the result of `make_const_list`; embedded execution and the DaCe + #: lowering use `common.ConstList` for the same thing. + #: TODO(egparedes): use `common.ConstList` in type inference too, and drop `None`. offset_type: common.Dimension | None diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index 500a355dbb..c6ca9f58ea 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -917,3 +917,12 @@ def test_a_dimension_cannot_be_staggered_twice(self): def test_resolve_rejects_a_bracketed_tag_of_another_owner(self): with pytest.raises(ValueError, match="not a parametrized dimension"): common.resolve(f"{KDim.tag}[{KDim.tag}]") + + +def test_resolve_loaded(): + # `resolve_loaded` never imports: it answers for what is loaded and gives up otherwise + assert common.resolve_loaded(IDim.tag) is IDim + assert common.resolve_loaded(common.Staggered[IDim].tag) is common.Staggered[IDim] + assert common.resolve_loaded("not_imported_anywhere.IDim") is None + assert common.resolve_loaded(f"{__name__}.does_not_exist") is None + assert common.resolve_loaded(f"{__name__}.test_resolve_loaded") is None # not a dimension From c8a69f0d36762a76c2a2df2e414c5e5ae0a18d8c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 00:03:45 +0200 Subject: [PATCH 17/20] docs[next]: renumber the dimensions-as-nominal-types ADR to 0029 ADR 0028 on main is now 'Plain Builders Instead of factory-boy Factories' (#2808), so the two ADRs of this stack move up by one. --- .../ADRs/next/0026-Staggered_Dimensions.md | 4 ++-- ..._Types.md => 0029-Dimensions_As_Nominal_Types.md} | 0 docs/development/ADRs/next/README.md | 2 +- src/gt4py/next/common.py | 12 ++++++------ src/gt4py/next/fingerprinting.py | 2 +- src/gt4py/next/iterator/ir.py | 2 +- src/gt4py/next/iterator/pretty_parser.py | 2 +- .../next/iterator/transforms/fuse_as_fieldop.py | 2 +- .../iterator/transforms/prune_empty_concat_where.py | 2 +- .../next/iterator/transforms/remove_broadcast.py | 2 +- .../next/iterator/type_system/type_synthesizer.py | 2 +- src/gt4py/next/otf/runners.py | 2 +- .../runners/dace/lowering/gtir_to_sdfg_lambda.py | 2 +- .../next/program_processors/runners/roundtrip.py | 2 +- src/gt4py/next/type_system/mypy_plugin.py | 2 +- tests/next_tests/fixtures/past_common.py | 2 +- tests/next_tests/integration_tests/cases_utils.py | 2 +- .../feature_tests/dace_tests/test_orchestration.py | 4 ++-- .../multi_feature_tests/fvm_nabla_setup.py | 2 +- .../iterator_tests/test_fvm_nabla.py | 2 +- .../unit_tests/embedded_tests/test_nd_array_field.py | 2 +- .../unit_tests/ffront_tests/test_type_deduction.py | 2 +- .../transforms_tests/test_unroll_reduce.py | 2 +- .../build_systems_tests/conftest.py | 2 +- .../next_tests/unit_tests/otf_tests/test_runners.py | 2 +- .../runners_tests/dace_tests/test_dace_bindings.py | 2 +- tests/next_tests/unit_tests/test_common.py | 2 +- 27 files changed, 33 insertions(+), 33 deletions(-) rename docs/development/ADRs/next/{0028-Dimensions_As_Nominal_Types.md => 0029-Dimensions_As_Nominal_Types.md} (100%) diff --git a/docs/development/ADRs/next/0026-Staggered_Dimensions.md b/docs/development/ADRs/next/0026-Staggered_Dimensions.md index 9cde0a7179..1f27b5fe18 100644 --- a/docs/development/ADRs/next/0026-Staggered_Dimensions.md +++ b/docs/development/ADRs/next/0026-Staggered_Dimensions.md @@ -10,7 +10,7 @@ tags: [] - **Updated**: 2026-09-23 > The *encoding* of this record is superseded by -> [ADR 0028](0028-Dimensions_As_Nominal_Types.md): a staggered dimension is the +> [ADR 0029](0029-Dimensions_As_Nominal_Types.md): a staggered dimension is the > real, interned class `Staggered[D]`, not a name prefix. The semantics below -- > half-integer positions, the shift convention, and the gtfn and DaCe treatment > -- are unchanged. @@ -66,7 +66,7 @@ index arithmetic is encoded in `common.connectivity_for_cartesian_shift`. A staggered dimension was encoded as its base dimension's name with the internal `_Staggered` prefix, rather than as a new attribute on `Dimension`. Since -[ADR 0028](0028-Dimensions_As_Nominal_Types.md) it is the class +[ADR 0029](0029-Dimensions_As_Nominal_Types.md) it is the class `common.Staggered[D]`, whose identity is the class itself; the helpers `is_staggered`, `flip_staggered` and `as_non_staggered` are unchanged in meaning and now read the class's `base`. diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md similarity index 100% rename from docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md rename to docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index fd74e97eb4..4f1661d745 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -22,7 +22,7 @@ Writing a new ADR is simple: - [0021 - Argument Descriptors](0021-Argument-Descriptors.md) - [0023 - Fingerprinting](0023-Fingerprinting.md) - [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md) -- [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md) +- [0029 - Dimensions as Nominal Types](0029-Dimensions_As_Nominal_Types.md) ### Frontend and Parsing #frontend diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 5839d2bbe0..c936f45d84 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -211,7 +211,7 @@ def __eq__( # type: ignore[misc] ) -> bool | Domain: # NOTE: dimension-vs-dimension comparison is deliberately *not* handled here. A # dimension's identity is its type, so `type.__eq__` (identity) is the correct - # answer; overriding it with `(tag, kind)` equality is what ADR 0028 rejects. + # answer; overriding it with `(tag, kind)` equality is what ADR 0029 rejects. if isinstance(value, DimensionMeta): return NotImplemented # both sides decline, so Python falls back to identity if isinstance(value, core_defs.INTEGRAL_TYPES): @@ -292,7 +292,7 @@ def __repr__(self) -> str: def __eq__(self, other: object) -> bool: if isinstance(other, DimensionIndex): - # NOTE: `is`, not `==`: a dimension's identity is its type (ADR 0028). + # NOTE: `is`, not `==`: a dimension's identity is its type (ADR 0029). return type(self) is type(other) and self.value == other.value return NotImplemented @@ -316,7 +316,7 @@ def dim(self) -> Dimension: #: The cost is that `get_origin()` of a PEP 695 alias is `None` rather than the aliased #: origin, so a site dispatching on an annotation's shape must resolve it first (see #: `xtyping.resolve_annotation`). `eve.datamodels` stores annotations *unresolved*, so this -#: applies to anything reading `__datamodel_fields__[...].type` too. See #2841 and ADR 0028. +#: applies to anything reading `__datamodel_fields__[...].type` too. See #2841 and ADR 0029. type Dimension = type[DimensionIndex] @@ -1003,7 +1003,7 @@ def __gt_domain__(self) -> Domain: def __gt_dims__(self) -> tuple[str, ...]: # NOTE: the unqualified name, not the `tag`. This is the interop protocol with # `gt4py.cartesian`, which identifies axes by their bare names (`"I"`, `"J"`, `"K"`); a - # qualified tag would not match and the axes would be transposed wrongly (ADR 0028). + # qualified tag would not match and the axes would be transposed wrongly (ADR 0029). return tuple(d.__qualname__ for d in self.__gt_domain__.dims) @@ -1712,7 +1712,7 @@ def __gt_builtin_func__(cls, /, func: fbuiltins.BuiltInFunction[_R, _P]) -> Call _DEFAULT_SKIP_VALUE: Final[int] = -1 #: Interned staggered dimensions, keyed by their *base dimension class*. #: -#: NOTE: this is not the name-keyed dimension registry ADR 0028 rejects. It is memoization of +#: NOTE: this is not the name-keyed dimension registry ADR 0029 rejects. It is memoization of #: a type constructor -- keyed by identity, populated only by `StaggeredMeta.__getitem__`, and #: never consulted to turn a user-authored name into a class. `typing`'s own subscription cache #: plays the same role for generic aliases. @@ -1725,7 +1725,7 @@ class StaggeredMeta(DimensionMeta): A PEP 695 generic cannot be used here: `Staggered[K]` would be a `typing._GenericAlias`, not a class, so it would fail `issubclass` and eve's `type[DimensionIndex]` validation, - and its `tag` could not name the base dimension. See ADR 0028. + and its `tag` could not name the base dimension. See ADR 0029. """ #: Set by `__getitem__` on each parametrization. Its presence is what distinguishes a diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index b762da1ec9..27316eaac5 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -225,7 +225,7 @@ def _instance_state(obj: Any) -> tuple[Any, ...]: # A parametrized dimension such as `Staggered[K]` has no importable qualified name (its # `__qualname__` contains brackets), so the by-reference `type` deconstruction rejects it. # It is fully determined by its base dimension, which *is* importable -- the same reduction - # its `copyreg` registration uses. The bare `Staggered` base is an ordinary class. See ADR 0028. + # its `copyreg` registration uses. The bare `Staggered` base is an ordinary class. See ADR 0029. common.StaggeredMeta: lambda obj: ( Deconstruction.from_pieces(obj.base, state=b"staggered_dimension") if "base" in obj.__dict__ diff --git a/src/gt4py/next/iterator/ir.py b/src/gt4py/next/iterator/ir.py index 05561c2398..717188c8a9 100644 --- a/src/gt4py/next/iterator/ir.py +++ b/src/gt4py/next/iterator/ir.py @@ -90,7 +90,7 @@ class OffsetLiteral(Expr): class AxisLiteral(Expr): - #: The dimension's tag, its qualified Python name (ADR 0028). + #: The dimension's tag, its qualified Python name (ADR 0029). value: str @property diff --git a/src/gt4py/next/iterator/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index 0324499d26..84a8008bbf 100644 --- a/src/gt4py/next/iterator/pretty_parser.py +++ b/src/gt4py/next/iterator/pretty_parser.py @@ -35,7 +35,7 @@ TYPE_LITERAL: CNAME INT_LITERAL: SIGNED_INT FLOAT_LITERAL: SIGNED_FLOAT - // A dimension or offset tag is a qualified Python name (ADR 0028): dotted, and -- for a + // A dimension or offset tag is a qualified Python name (ADR 0029): dotted, and -- for a // parametrized dimension such as `Staggered[pkg.K]` -- with one bracketed dotted name. // Unambiguous here: a tag starts with a letter (a float does not), and the literal's // suffix terminates it. diff --git a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py index 4bbc34e0f2..de749a356b 100644 --- a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py +++ b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py @@ -252,7 +252,7 @@ class FuseAsFieldOp( >>> from gt4py.next import utils >>> from gt4py.next.iterator.ir_utils import ir_makers as im >>> class IDim(gtx.DimensionIndex): ... - >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: >>> import sys >>> sys.modules[__name__].IDim = IDim 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 2bc8e84cc4..e0db04088a 100644 --- a/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py +++ b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py @@ -85,7 +85,7 @@ class _PruneEmptyConcatWhere(PreserveLocationVisitor, NodeTranslator): >>> from gt4py.next import common >>> class IDim(common.DimensionIndex): ... - >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: >>> import sys >>> sys.modules[__name__].IDim = IDim diff --git a/src/gt4py/next/iterator/transforms/remove_broadcast.py b/src/gt4py/next/iterator/transforms/remove_broadcast.py index 1a92361216..f3dcb8a018 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -26,7 +26,7 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): >>> from gt4py.next import Dimension, common >>> from gt4py.next.common import DimensionIndex >>> class IDim(DimensionIndex): ... - >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: >>> class JDim(DimensionIndex): ... >>> import sys diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index bcda84959e..098c403755 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -511,7 +511,7 @@ def _resolve_dimensions( >>> class IDim(common.DimensionIndex): ... >>> IHalfDim = common.flip_staggered(IDim) >>> class JDim(common.DimensionIndex): ... - >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0028), and a + >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: >>> import sys >>> sys.modules[__name__].Edge = Edge diff --git a/src/gt4py/next/otf/runners.py b/src/gt4py/next/otf/runners.py index 8ee324c6d5..7057e73137 100644 --- a/src/gt4py/next/otf/runners.py +++ b/src/gt4py/next/otf/runners.py @@ -189,7 +189,7 @@ def _interactive_main_reference(obj: object) -> str | None: there resolve in the worker. An interactive `__main__` -- a notebook kernel, the REPL, `python -c` -- has no `__file__` to re-import, so a class declared there pickles fine in the parent (which has it) and then fails to unpickle in the worker. Dimensions are classes - identified by their qualified name (ADR 0028), which makes this the common case in + identified by their qualified name (ADR 0029), which makes this the common case in notebooks. Only scans when `__main__` is interactive, so ordinary scripts pay nothing. 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 b5a3a1d632..a28aad41c3 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 @@ -1152,7 +1152,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: # NOTE: the connectivity's own local dimension, not one synthesized from the offset # tag. The latter named a local dimension after the *offset*, which only coincided with # the real one under the old `V2EDim = Dimension("V2E")` convention, and under nominal - # identity (ADR 0028) a tag string cannot be turned back into a dimension at all. + # identity (ADR 0029) a tag string cannot be turned back into a dimension at all. offset_type = conn_type.neighbor_dim neighbor_idx = gtir_to_sdfg_utils.get_map_variable(offset_type) diff --git a/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 9bb62bd5de..90e5983527 100644 --- a/src/gt4py/next/program_processors/runners/roundtrip.py +++ b/src/gt4py/next/program_processors/runners/roundtrip.py @@ -185,7 +185,7 @@ def _generate_source( f'{common.codegen_name(o)} = offset("{o}")' for o in offset_literals ) # A dimension is not constructed from its name any more: its tag is its qualified Python - # name, so the emitted program imports it (ADR 0028). + # name, so the emitted program imports it (ADR 0029). axis_literals_src = "\n".join( f'{common.codegen_name(o.value)} = gtx.resolve("{o.value}")' for o in axis_literals_set ) diff --git a/src/gt4py/next/type_system/mypy_plugin.py b/src/gt4py/next/type_system/mypy_plugin.py index d1c423e325..35b90168e0 100644 --- a/src/gt4py/next/type_system/mypy_plugin.py +++ b/src/gt4py/next/type_system/mypy_plugin.py @@ -32,7 +32,7 @@ Dimensions no longer need plugin support: a concrete dimension is a class ('class IDim(gtx.DimensionIndex): ...'), which is a valid annotation for any type checker. See ADR -0028. Only the mixed-precision hooks below remain; this plugin is scheduled for removal once +0029. Only the mixed-precision hooks below remain; this plugin is scheduled for removal once dtype-generic fields land. """ diff --git a/tests/next_tests/fixtures/past_common.py b/tests/next_tests/fixtures/past_common.py index 43434f36ed..65a834771f 100644 --- a/tests/next_tests/fixtures/past_common.py +++ b/tests/next_tests/fixtures/past_common.py @@ -13,7 +13,7 @@ import gt4py.next as gtx from gt4py.next import float64 -# NOTE: imported, not redeclared. Under nominal identity (ADR 0028) a same-named declaration +# NOTE: imported, not redeclared. Under nominal identity (ADR 0029) a same-named declaration # here would be a different dimension from the one `cases_utils` declares, where the old # `Dimension("...")` values compared equal -- and tests mix objects from both modules. from next_tests.integration_tests.cases_utils import IDim diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 464bf7fa1f..adc52ee8e2 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -33,7 +33,7 @@ # NOTE: the unstructured dimensions are declared once, in `toy_connectivity`, and imported here. # Both modules used to declare their own `Dimension("Vertex")` etc., which compared equal; under -# nominal identity (ADR 0028) that would be two different dimensions, and tests that mix a +# nominal identity (ADR 0029) that would be two different dimensions, and tests that mix a # `toy_connectivity` connectivity with a `cases_utils` mesh would silently stop matching. from next_tests.toy_connectivity import C2EDim, Cell, E2VDim, Edge, V2EDim, Vertex diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py index 47551e473c..7b711a256d 100644 --- a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py +++ b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_orchestration.py @@ -150,7 +150,7 @@ def get_stride_from_numpy_to_dace(arg: core_defs.NDArrayObject, axis: int) -> in offset_provider, rows=3, cols=2, - # the connectivity argument is named after the mangled offset key (ADR 0028) + # the connectivity argument is named after the mangled offset key (ADR 0029) **{E2V_CONN: e2v}, **{f"__{E2V_CONN}_source_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 0)}, **{f"__{E2V_CONN}_neighbor_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 1)}, @@ -172,7 +172,7 @@ def get_stride_from_numpy_to_dace(arg: core_defs.NDArrayObject, axis: int) -> in offset_provider, rows=3, cols=2, - # the connectivity argument is named after the mangled offset key (ADR 0028) + # the connectivity argument is named after the mangled offset key (ADR 0029) **{E2V_CONN: e2v}, **{f"__{E2V_CONN}_source_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 0)}, **{f"__{E2V_CONN}_neighbor_stride": get_stride_from_numpy_to_dace(e2v.ndarray, 1)}, diff --git a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py index 4fbb01c72d..99d8085dd0 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py @@ -36,7 +36,7 @@ from gt4py import next as gtx from gt4py.next.iterator import atlas_utils -# NOTE: imported, not redeclared. Under nominal identity (ADR 0028) a same-named declaration +# NOTE: imported, not redeclared. Under nominal identity (ADR 0029) a same-named declaration # here would be a different dimension from the one `toy_connectivity` declares, where the old # `Dimension("...")` values compared equal -- and tests mix objects from both modules. from next_tests.toy_connectivity import E2VDim, Edge, V2EDim, Vertex diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py index 3697652641..335079d1a4 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py @@ -31,7 +31,7 @@ from gt4py.next.iterator.runtime import set_at, fendef, fundef, offset # NOTE: the dimensions are imported, not redeclared. The connectivities come from -# `nabla_setup` and are built on *its* dimension classes; under nominal identity (ADR 0028) a +# `nabla_setup` and are built on *its* dimension classes; under nominal identity (ADR 0029) a # same-named redeclaration here would be a different dimension, where it used to compare equal. from next_tests.integration_tests.multi_feature_tests.fvm_nabla_setup import ( E2VDim, diff --git a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py index bef6a29c8a..c7b184dcff 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 @@ -146,7 +146,7 @@ def _default_dim(i: int) -> Dimension: Return the `i`-th default dimension, `D`, creating it on first use. The dimensions must be *stable across calls*: two domains built by separate calls are - compared by the tests, and under nominal identity (ADR 0028) a fresh class per call would + compared by the tests, and under nominal identity (ADR 0029) a fresh class per call would be a different dimension. Each is bound as a module attribute, so it is also importable by its qualified name like any other dimension. """ diff --git a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py index 88a8c640d8..63fa61893d 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 @@ -39,7 +39,7 @@ from next_tests.artifacts import custom_named_collections as cnc # NOTE: the named collections in the artifact are annotated with *its* `TDim`, and the expected -# types below are compared against them. Under nominal identity (ADR 0028) a redeclared `TDim` +# types below are compared against them. Under nominal identity (ADR 0029) a redeclared `TDim` # here would be a different dimension, where the old `Dimension("TDim")` compared equal. TDim = cnc.TDim diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py index 5451a2dc44..1028fd987f 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 @@ -26,7 +26,7 @@ class dummy_neighbor(common.DimensionIndex): ... #: The local dimensions of the neighbor lists under test. Each one's `tag` is also its IR offset #: string and its offset-provider key: `UnrollReduce` looks a connectivity up by the local -#: dimension of the list it reduces, so those three names must be a single string (ADR 0028). +#: 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): ... diff --git a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py index 00c77dfd58..7160b12de4 100644 --- a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py +++ b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py @@ -56,7 +56,7 @@ def make_program_source(name: str) -> artifacts.ProgramSource: returns=True, ) # NOTE: the tag types are named after the *mangled* dimension tags, which is what the - # generated bindings reference; a dimension's tag is its qualified Python name (ADR 0028). + # generated bindings reference; a dimension's tag is its qualified Python name (ADR 0029). i_t, j_t = (common.codegen_name(d.tag) for d in (I, J)) func = cpp_interface.render_function_declaration( entry_point, diff --git a/tests/next_tests/unit_tests/otf_tests/test_runners.py b/tests/next_tests/unit_tests/otf_tests/test_runners.py index 2ea966ff7e..0a0941d792 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_runners.py +++ b/tests/next_tests/unit_tests/otf_tests/test_runners.py @@ -352,7 +352,7 @@ class TestInteractiveMainReference: """ A class declared in an interactive `__main__` pickles in the parent but not in a worker. - Dimensions are classes identified by their qualified name (ADR 0028), so a notebook that + Dimensions are classes identified by their qualified name (ADR 0029), so a notebook that declares one would otherwise break the default process-pool compilation. """ diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index 6ecaa45fcc..ff1ec0addd 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -23,7 +23,7 @@ from next_tests.integration_tests import cases, cases_utils # NOTE: from `cases`, not `test_common`: the `cartesian_case` fixture is sized on the `cases` -# dimensions, and under nominal identity (ADR 0028) another module's same-named `IDim` is a +# dimensions, and under nominal identity (ADR 0029) another module's same-named `IDim` is a # different dimension -- it used to compare equal. from next_tests.integration_tests.cases import E2V, V2E, E2VDim, IDim, JDim, KDim, V2EDim diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index c6ca9f58ea..a2c4832b78 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -895,7 +895,7 @@ def test_injective_and_reversible_exhaustively(self): def test_gt_dims_are_unqualified_names(): """ `__gt_dims__` is the interop protocol with `gt4py.cartesian`, which names axes by their bare - names (`"I"`, `"J"`, `"K"`). A dimension's `tag` is its qualified name (ADR 0028), which + names (`"I"`, `"J"`, `"K"`). A dimension's `tag` is its qualified name (ADR 0029), which cartesian would not recognize, and would then transpose the array wrongly. """ field = gtx.as_field([IDim, JDim], np.zeros((2, 3))) From 47b276bfc0a3352bd4d6e2d753f27653991bf875 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 00:20:07 +0200 Subject: [PATCH 18/20] feat[next]: Cartesian axis levels below DimensionIndex; bound Staggered on a declared axis Add `AnyCartesianAxisIndex` (either cell class of a Cartesian axis) and `CartesianAxisIndex` (a declared axis) below `DimensionIndex`, both exported as `gtx.*`. `Staggered[D]` derives from `AnyCartesianAxisIndex` and its parameter is bounded on `CartesianAxisIndex`, so a doubly staggered dimension, a staggered mesh location and a staggered local dimension are `[type-var]` errors for mypy and pyright and `TypeError`s at runtime. `DimensionMeta.__add__` / `__sub__` carry the self-type `type[AnyCartesianAxisIndex]`, so `Cell + 1` on a mesh location is an `[operator]` error, a `TypeError` at runtime and a `DSLError` in a field operator. Comparisons (`D == n`, `D < n`) stay available on every dimension. The tree's Cartesian dimensions (`IDim`, `KDim`, ...) are declared as `CartesianAxisIndex`; mesh locations stay `DimensionIndex`. ADR 0029 records the levels and the equality decision (dimension-vs-dimension `==` is identity). --- .../ADRs/next/0026-Staggered_Dimensions.md | 6 +- .../next/0029-Dimensions_As_Nominal_Types.md | 94 +++++++++++++-- docs/user/next/QuickstartGuide.md | 8 +- .../exercises/1_simple_addition.ipynb | 4 +- .../1_simple_addition_solution.ipynb | 4 +- docs/user/next/workshop/exercises/helpers.py | 4 +- docs/user/next/workshop/slides/slides_1.ipynb | 2 +- docs/user/next/workshop/slides/slides_2.ipynb | 2 +- docs/user/next/workshop/slides/slides_3.ipynb | 2 +- docs/user/next/workshop/slides/slides_4.ipynb | 2 +- examples/lap_cartesian_vs_next.ipynb | 6 +- src/gt4py/next/__init__.py | 4 + src/gt4py/next/common.py | 107 ++++++++++++++---- src/gt4py/next/constructors.py | 14 +-- src/gt4py/next/embedded/common.py | 6 +- src/gt4py/next/ffront/decorator.py | 2 +- .../ffront/foast_passes/type_deduction.py | 18 ++- src/gt4py/next/ffront/foast_pretty_printer.py | 4 +- src/gt4py/next/ffront/foast_to_gtir.py | 4 +- src/gt4py/next/ffront/foast_to_past.py | 2 +- src/gt4py/next/ffront/func_to_foast.py | 6 +- src/gt4py/next/ffront/func_to_past.py | 2 +- src/gt4py/next/ffront/past_to_itir.py | 6 +- src/gt4py/next/ffront/type_info.py | 2 +- src/gt4py/next/field_utils.py | 4 +- src/gt4py/next/iterator/embedded.py | 4 +- src/gt4py/next/iterator/ir_utils/ir_makers.py | 6 +- .../iterator/transforms/fuse_as_fieldop.py | 2 +- .../iterator/transforms/inline_fundefs.py | 2 +- .../transforms/prune_empty_concat_where.py | 2 +- .../iterator/transforms/remove_broadcast.py | 6 +- ...replace_get_domain_range_with_constants.py | 2 +- .../iterator/type_system/type_synthesizer.py | 6 +- src/gt4py/next/type_system/mypy_plugin.py | 2 +- src/gt4py/next/type_system/type_info.py | 16 +-- .../artifacts/custom_named_collections.py | 3 +- .../benchmarks/benchmark_program_call.py | 2 +- .../integration_tests/cases_utils.py | 6 +- ..._write_back_buffer_elimination_lowering.py | 2 +- .../ffront_tests/test_foast_pretty_printer.py | 5 +- .../instrumentation_tests/test_hooks.py | 2 +- .../iterator_tests/test_builtins.py | 2 +- .../iterator_tests/test_conditional.py | 2 +- .../iterator_tests/test_implicit_fencil.py | 2 +- .../iterator_tests/test_program.py | 2 +- .../iterator_tests/test_tuple.py | 6 +- .../iterator_tests/test_anton_toy.py | 6 +- .../iterator_tests/test_if_stmt.py | 2 +- .../iterator_tests/test_temporaries.py | 4 +- .../embedded_tests/test_domain_pickle.py | 4 +- .../embedded_tests/test_basic_program.py | 2 +- .../unit_tests/embedded_tests/test_common.py | 6 +- .../unit_tests/embedded_tests/test_context.py | 4 +- .../embedded_tests/test_nd_array_field.py | 19 ++-- .../test_decorator_domain_deduction.py | 4 +- .../ffront_tests/test_diagnostic_messages.py | 2 +- .../unit_tests/ffront_tests/test_fbuiltins.py | 2 +- .../ffront_tests/test_foast_to_gtir.py | 4 +- .../ffront_tests/test_func_to_foast.py | 6 +- .../test_func_to_foast_error_line_number.py | 2 +- .../unit_tests/ffront_tests/test_stages.py | 2 +- .../ffront_tests/test_type_deduction.py | 26 +++-- .../ir_utils_tests/test_domain_utils.py | 6 +- .../iterator_tests/test_embedded_internals.py | 2 +- .../test_inline_dynamic_shifts.py | 2 +- .../iterator_tests/test_pretty_printer.py | 4 +- .../iterator_tests/test_runtime_domain.py | 2 +- ...t_concat_where_canonicalize_domain_args.py | 2 +- .../test_concat_where_expand_tuple_args.py | 2 +- ...st_concat_where_transform_to_as_fieldop.py | 4 +- .../transforms_tests/test_cse.py | 2 +- .../test_dead_code_elimination.py | 2 +- .../transforms_tests/test_domain_inference.py | 6 +- .../test_expand_tuple_maps.py | 2 +- .../transforms_tests/test_fuse_as_fieldop.py | 4 +- .../transforms_tests/test_global_tmps.py | 6 +- .../transforms_tests/test_inline_scalar.py | 2 +- .../transforms_tests/test_prune_casts.py | 2 +- .../test_prune_empty_concat_where.py | 2 +- .../build_systems_tests/conftest.py | 4 +- .../otf_tests/test_compiled_program.py | 2 +- .../gtfn_tests/test_gtfn_module.py | 2 +- .../transformation_tests/test_map_promoter.py | 4 +- tests/next_tests/unit_tests/test_common.py | 64 +++++++++-- .../unit_tests/test_constructors.py | 6 +- .../test_custom_layout_allocators.py | 14 +-- .../next_tests/unit_tests/test_field_utils.py | 2 +- tests/next_tests/unit_tests/test_utils.py | 4 +- .../type_system_tests/test_type_info.py | 9 +- .../test_type_translation.py | 4 +- typing_tests/test_next.yaml | 57 +++++++--- 91 files changed, 482 insertions(+), 234 deletions(-) diff --git a/docs/development/ADRs/next/0026-Staggered_Dimensions.md b/docs/development/ADRs/next/0026-Staggered_Dimensions.md index 1f27b5fe18..7584508c8d 100644 --- a/docs/development/ADRs/next/0026-Staggered_Dimensions.md +++ b/docs/development/ADRs/next/0026-Staggered_Dimensions.md @@ -7,7 +7,7 @@ tags: [] - **Status**: valid - **Authors**: Till Ehrengruber (@tehrengruber) - **Created**: 2026-07-08 -- **Updated**: 2026-09-23 +- **Updated**: 2026-10-02 > The *encoding* of this record is superseded by > [ADR 0029](0029-Dimensions_As_Nominal_Types.md): a staggered dimension is the @@ -69,7 +69,9 @@ A staggered dimension was encoded as its base dimension's name with the internal [ADR 0029](0029-Dimensions_As_Nominal_Types.md) it is the class `common.Staggered[D]`, whose identity is the class itself; the helpers `is_staggered`, `flip_staggered` and `as_non_staggered` are unchanged in meaning -and now read the class's `base`. +and now read the class's `base`. Only a declared Cartesian axis +(`CartesianAxisIndex`) can be staggered, which makes a doubly staggered dimension +and a staggered mesh location type errors. Dimensions are identified by their **name** and appear throughout the toolchain in more than one form: as a `common.Dimension` instance, but also as an diff --git a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md index d3611694f4..f1ec832a45 100644 --- a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md @@ -7,17 +7,20 @@ tags: [] - **Status**: proposed - **Authors**: Enrique González Paredes (@egparedes) - **Created**: 2026-09-18 -- **Updated**: 2026-09-24 +- **Updated**: 2026-10-02 A concrete dimension becomes a **class**, and an index along it an **instance** of that class — the shape `enum.Enum` uses, where the class is the collection and the instances are its members: ```python -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... -class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... +class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class Cell(gtx.DimensionIndex): ... # a mesh location: not a Cartesian axis IDim # the dimension -- annotated `gtx.Dimension` @@ -29,7 +32,9 @@ no gt4py mypy plugin. A dimension's **identity is the Python type**, and its `tag` — the string that crosses into the IR and the generated code — is its **qualified Python name**, -`f"{cls.__module__}.{cls.__qualname__}"`. +`f"{cls.__module__}.{cls.__qualname__}"`. A Cartesian axis — a dimension with +index arithmetic and a staggered partner — is a subclass of `CartesianAxisIndex`; +mesh locations subclass `DimensionIndex` directly. ## Context @@ -157,6 +162,8 @@ disappears. - bases are `(Staggered,)` and deliberately **not** `(Staggered, base)`: a staggered dimension is a *different* dimension, so `issubclass(Staggered[KDim], KDim)` must be false. Only `kind` is inherited. + `Staggered` itself derives from `AnyCartesianAxisIndex`, and its parameter is + bounded on `CartesianAxisIndex` (see *Cartesian axes* below). - `is_staggered(dim)` is `"base" in dim.__dict__` and `as_non_staggered(dim)` is `dim.base` — structural, no string sniffing. `issubclass(dim, Staggered)` would be wrong, because it is also true of the @@ -173,6 +180,48 @@ disappears. staggered metaclass. It must fall back to by-reference pickling for the bare base, which is also an instance of that metaclass. +### Cartesian axes are a level of the hierarchy + +A **Cartesian axis** is an index space with integer index arithmetic and exactly one +staggered partner — the dimensions `CartesianConnectivity` acts on. One axis of a +Cartesian grid is a 1-dimensional cell complex with exactly two cell classes; a +declared axis and its `Staggered[...]` name those two, and `Staggered` is the +involution that swaps them. Mesh locations are not axes: an unstructured mesh does +not factor into per-axis cell classes, so it has no half cells and no index +arithmetic. Two levels below the root encode this: + +``` +DimensionIndex # the root; every `type[DimensionIndex]` keeps its meaning +├── AnyCartesianAxisIndex # either cell class of a Cartesian axis +│ ├── CartesianAxisIndex # a declared axis: what users subclass +│ └── Staggered[D: CartesianAxisIndex] # its derived partner +└── (direct subclasses) # mesh locations, and index spaces without geometry +``` + +Both levels sit *below* `DimensionIndex`, so `Staggered[K]` stays a `DimensionIndex` +and no annotation or `issubclass` guard in the tree widens. The bound is on the +*declared* level, and `Staggered[K]` is only an `AnyCartesianAxisIndex`, so: + +| Rejected | Statically (mypy, pyright) | At runtime | +| --- | --- | --- | +| `Staggered[Staggered[K]]` | `[type-var]` | `TypeError` | +| `Staggered[Cell]`, staggering a local dimension | `[type-var]` | `TypeError` | +| `Cell + 1`, `Cell - 1` | `[operator]` | `TypeError`; a `DSLError` in a field operator | + +The last row needs `DimensionMeta.__add__` / `__sub__` declared with the self-type +`cls: type[AnyCartesianAxisIndex]`. Both checkers bind it correctly at every call +site and both reject it at the definition site, with different diagnostics (mypy +`[misc]`, pyright `reportGeneralTypeIssues`), so it costs two separately spelled +suppressions. The runtime check covers unannotated code; hand-written iterator IR, +which names dimensions by tag, is not checked. Comparisons are deliberately *not* restricted: `D == n` +and `D < n` build a `Domain` on every dimension, as `concat_where` over a mesh +location requires. + +Whether a dimension is an axis or a mesh location is a decision per declaration — +`IDim` and `Cell` are both `HORIZONTAL`, and only the first is an axis — so it cannot +be derived from `kind`. `CartesianAxisIndex` and `AnyCartesianAxisIndex` are both +exported as `gtx.*`. + ### `Dimension` is annotation-only `common.Dimension` becomes a PEP 695 alias for `type[DimensionIndex]`, so the @@ -216,10 +265,22 @@ declaration would silently reorder a field's dimensions. It is therefore keyed o the unqualified name, with `tag` only breaking ties between same-named dimensions from different modules so the order stays total. +### Equality between dimensions is identity + +`DimensionMeta.__eq__` stays, for the `I == 5` → `Domain` overload that +`concat_where` uses, but it does not compare dimensions: for a dimension operand it +returns `NotImplemented`, the reflected call does the same, and Python falls back to +identity. So `I == J` is always a `bool`, and it is `True` only for the same class — +"equality is `is`" holds for every dimension operand. The one non-`bool` result is +`dim == `, which builds a `Domain`, and `Domain.__bool__` raises. A dict +lookup reaches that path only if an integer key and a dimension-class key share a +hash bucket in one dict; class-keyed mappings in the tree (domains, offset +providers) mix classes with strings at most, and `str` against a class compares +`False`. That residual hazard is accepted. + `DimensionMeta` must declare `__hash__ = type.__hash__` explicitly: Python sets -`__hash__ = None` on any class body defining `__eq__` without it, and `__eq__` -stays for the `I == 5` → `Domain` overload. Without it every dimension class is -unhashable and `ts.DimensionType` fails at *import*. +`__hash__ = None` on any class body defining `__eq__` without it. Without it every +dimension class is unhashable and `ts.DimensionType` fails at *import*. ## Consequences @@ -249,6 +310,20 @@ unhashable and `ts.DimensionType` fails at *import*. remove `resolve` from the type-inference hot path, but changes IR node shape in the same step as the dimension rewrite. Deferred. - **`Staggered` as a PEP 695 generic.** Does not produce a class; see Decision 6. +- **A sibling root above `DimensionIndex`** for the staggered dimensions. Gives the + same static ban on double staggering, but `Staggered[K]` stops being a + `DimensionIndex`, so every annotation and `issubclass` guard that must accept a + staggered dimension widens. The axis levels sit below the root instead. +- **A structural discriminator** (a `Protocol` with a `Literal[False]` class + variable the staggered class overrides). Needs a suppressed incompatible + override, gives an opaque diagnostic, and buys only the double-staggering ban: + with no axis concept, `Cell + 1` and `Staggered[Cell]` stay unchecked. +- **`Staggered[D: AnyCartesianAxisIndex]`**, bounding on either cell class. Readmits + `Staggered[Staggered[K]]`; the bound has to name the declared level. +- **A second partner constructor** for the other alignment (`i + 1/2` instead of + `i - 1/2`). Gives three cell classes per axis where an axis has two and makes + `flip_staggered` partial. Either alignment is already reachable by choosing which + member of the pair to declare. ## References @@ -260,5 +335,6 @@ unhashable and `ts.DimensionType` fails at *import*. [ADR 0026](0026-Staggered_Dimensions.md); the indexing convention there is unchanged. - Consequence for the build cache of [ADR 0023](0023-Fingerprinting.md). -- An alternative to #2844, which implements the same class-shaped dimension with - value identity. +- The Cartesian axis levels follow the specification in gt4py_knowledge#36. +- An alternative to #2844 (closed, superseded), which implemented the same + class-shaped dimension with value identity. diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 1e1cbfc280..63d5c2dda3 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -51,11 +51,11 @@ from gt4py.next import float64, neighbor_sum, where, Dims #### Fields -Fields store data as a multi-dimensional array, and are defined over a set of named dimensions. The code snippet below defines two dimensions, `CellDim` and `KDim` -- a dimension is a class -- and creates the fields `a` and `b` over their cartesian product using the `gtx.as_field` helper function. The fields contain the values 2 for `a` and 3 for `b` for all entries. +Fields store data as a multi-dimensional array, and are defined over a set of named dimensions. The code snippet below defines two dimensions, `CellDim` and `KDim` -- a dimension is a class. `CellDim` is a mesh location and subclasses `gtx.DimensionIndex`; `KDim` is a Cartesian axis, which supports index arithmetic such as `KDim + 1` and staggering, and subclasses `gtx.CartesianAxisIndex`. The snippet then creates the fields `a` and `b` over their cartesian product using the `gtx.as_field` helper function. The fields contain the values 2 for `a` and 3 for `b` for all entries. ```{code-cell} ipython3 class CellDim(gtx.DimensionIndex): ... -class KDim(gtx.DimensionIndex): ... +class KDim(gtx.CartesianAxisIndex): ... num_cells = 5 num_layers = 6 @@ -70,8 +70,8 @@ b = gtx.as_field([CellDim, KDim], np.full(shape=grid_shape, fill_value=b_value, Additional numpy-equivalent constructors are available, namely `ones`, `zeros`, `empty`, `full`. These require domain, dtype, and allocator (e.g. a backend) specifications. ```{code-cell} ipython3 -class I(gtx.DimensionIndex): ... -class J(gtx.DimensionIndex): ... +class I(gtx.CartesianAxisIndex): ... +class J(gtx.CartesianAxisIndex): ... array_of_ones_numpy = np.ones((grid_shape[0], grid_shape[1])) field_of_ones = gtx.ones( diff --git a/docs/user/next/workshop/exercises/1_simple_addition.ipynb b/docs/user/next/workshop/exercises/1_simple_addition.ipynb index 64cd1c55c0..dc2895dc4c 100644 --- a/docs/user/next/workshop/exercises/1_simple_addition.ipynb +++ b/docs/user/next/workshop/exercises/1_simple_addition.ipynb @@ -36,10 +36,10 @@ "metadata": {}, "outputs": [], "source": [ - "class I(gtx.DimensionIndex): ...\n", + "class I(gtx.CartesianAxisIndex): ...\n", "\n", "\n", - "class J(gtx.DimensionIndex): ...\n", + "class J(gtx.CartesianAxisIndex): ...\n", "\n", "\n", "size = 10" diff --git a/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb b/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb index 40409ad8bc..3ec51bf5d7 100644 --- a/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb +++ b/docs/user/next/workshop/exercises/1_simple_addition_solution.ipynb @@ -51,10 +51,10 @@ "metadata": {}, "outputs": [], "source": [ - "class I(gtx.DimensionIndex): ...\n", + "class I(gtx.CartesianAxisIndex): ...\n", "\n", "\n", - "class J(gtx.DimensionIndex): ...\n", + "class J(gtx.CartesianAxisIndex): ...\n", "\n", "\n", "size = 10" diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index c398524538..82c0ef6a57 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -11,7 +11,7 @@ import gt4py.next as gtx from gt4py.next.iterator.embedded import MutableLocatedField from gt4py.next import neighbor_sum, where, Dims -from gt4py.next import Dimension, DimensionIndex, DimensionKind, FieldOffset +from gt4py.next import CartesianAxisIndex, Dimension, DimensionIndex, DimensionKind, FieldOffset from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( run_gtfn as gtfn_cpu, @@ -386,7 +386,7 @@ class V(DimensionIndex): ... class E(DimensionIndex): ... -class K(DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... +class K(CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... diff --git a/docs/user/next/workshop/slides/slides_1.ipynb b/docs/user/next/workshop/slides/slides_1.ipynb index 724162dc8f..643950aa9c 100644 --- a/docs/user/next/workshop/slides/slides_1.ipynb +++ b/docs/user/next/workshop/slides/slides_1.ipynb @@ -210,7 +210,7 @@ "class Cell(gtx.DimensionIndex): ...\n", "\n", "\n", - "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ...\n", + "class K(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ...\n", "\n", "\n", "domain = gtx.domain({Cell: 5, K: 6})\n", diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index db8f370abc..6a44a6eef7 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -57,7 +57,7 @@ "class Cell(gtx.DimensionIndex): ...\n", "\n", "\n", - "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ..." + "class K(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ..." ] }, { diff --git a/docs/user/next/workshop/slides/slides_3.ipynb b/docs/user/next/workshop/slides/slides_3.ipynb index 1d699e996f..a85937180f 100644 --- a/docs/user/next/workshop/slides/slides_3.ipynb +++ b/docs/user/next/workshop/slides/slides_3.ipynb @@ -57,7 +57,7 @@ "class Cell(gtx.DimensionIndex): ...\n", "\n", "\n", - "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ..." + "class K(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ..." ] }, { diff --git a/docs/user/next/workshop/slides/slides_4.ipynb b/docs/user/next/workshop/slides/slides_4.ipynb index 12870c30de..4567390c20 100644 --- a/docs/user/next/workshop/slides/slides_4.ipynb +++ b/docs/user/next/workshop/slides/slides_4.ipynb @@ -54,7 +54,7 @@ "metadata": {}, "outputs": [], "source": [ - "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ..." + "class K(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ..." ] }, { diff --git a/examples/lap_cartesian_vs_next.ipynb b/examples/lap_cartesian_vs_next.ipynb index 9413ee57a4..fd4a4338c8 100644 --- a/examples/lap_cartesian_vs_next.ipynb +++ b/examples/lap_cartesian_vs_next.ipynb @@ -67,13 +67,13 @@ "\n", "\n", "# Note: for gt4py.next, names don't matter, for gt4py.cartesian they have to be \"I\", \"J\", \"K\"\n", - "class I(gtx.DimensionIndex): ...\n", + "class I(gtx.CartesianAxisIndex): ...\n", "\n", "\n", - "class J(gtx.DimensionIndex): ...\n", + "class J(gtx.CartesianAxisIndex): ...\n", "\n", "\n", - "class K(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ...\n", + "class K(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ...\n", "\n", "\n", "domain = gtx.domain({I: nx, J: ny, K: nz})\n", diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index e665024d7d..b8a7bf5143 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -22,6 +22,8 @@ from .._core.definitions import CUPY_DEVICE_TYPE, Device, DeviceType, is_scalar_type from . import common, ffront, iterator, program_processors, typing from .common import ( + AnyCartesianAxisIndex, + CartesianAxisIndex, CartesianConnectivity, Connectivity, Dimension, @@ -119,6 +121,8 @@ # from common "Dimension", "DimensionIndex", + "AnyCartesianAxisIndex", + "CartesianAxisIndex", "DimensionKind", "Staggered", "resolve", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index c936f45d84..a5486c024f 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -184,10 +184,26 @@ def __str__(cls) -> str: # display name; `repr` carries the module and disambiguates when it matters. return f"{cls.__qualname__}[{cls.kind}]" - def __add__(cls: Dimension, offset: int | float) -> Connectivity: # type: ignore[misc] + # 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`), + # and both reject it at the definition site, each with its own diagnostic -- mypy `[misc]` + # ("self" parameter missing) and pyright `reportGeneralTypeIssues` ("must be a supertype of + # its class") -- hence the two suppressions. The runtime check covers unannotated code. + def __add__( # type: ignore[misc] + cls: type[AnyCartesianAxisIndex], # pyright: ignore[reportGeneralTypeIssues] + offset: int | float, + ) -> Connectivity: + if not issubclass(cls, AnyCartesianAxisIndex): + raise TypeError( + f"'{cls.__qualname__}' is not a Cartesian axis: only a dimension declared as" + " 'CartesianAxisIndex' (or its 'Staggered[...]' partner) has index arithmetic." + ) return connectivity_for_cartesian_shift(cls, offset) - def __sub__(cls: Dimension, offset: int | float) -> Connectivity: # type: ignore[misc] + def __sub__( # type: ignore[misc] + cls: type[AnyCartesianAxisIndex], # pyright: ignore[reportGeneralTypeIssues] + offset: int | float, + ) -> Connectivity: return cls + (-offset) def __gt__(cls: Dimension, value: core_defs.IntegralScalar) -> Domain: # type: ignore[misc] @@ -243,8 +259,8 @@ class DimensionIndex(metaclass=DimensionMeta): is how it is spelled in the IR. `value` is an index position along it. Examples: - >>> class I(DimensionIndex): ... - >>> class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... + >>> class I(CartesianAxisIndex): ... + >>> class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... >>> str(I), K.kind ('I[horizontal]', ) @@ -255,7 +271,7 @@ class DimensionIndex(metaclass=DimensionMeta): Two dimension classes are the same dimension only if they are the same class: - >>> class I2(DimensionIndex): ... + >>> class I2(CartesianAxisIndex): ... >>> I == I2 False """ @@ -320,6 +336,52 @@ def dim(self) -> Dimension: type Dimension = type[DimensionIndex] +class AnyCartesianAxisIndex(DimensionIndex): + """ + Either cell class of a Cartesian axis: a declared `CartesianAxisIndex` or its `Staggered[...]`. + + One axis of a Cartesian grid has exactly two cell classes, and a declared axis and its + staggered partner name them (ADR 0029). Only Cartesian shifts need this level -- `D + n`, + `D - n` and `as_offset` -- while comparisons (`D == n`, `D < n`, which build a `Domain`) stay + available on every dimension, mesh locations included. + + Annotate with this level where any cell class of an axis is accepted; declare axes with + `CartesianAxisIndex`. + """ + + __slots__ = () + + +class CartesianAxisIndex(AnyCartesianAxisIndex): + """ + A declared Cartesian axis: integer index arithmetic and exactly one staggered partner. + + Subclass it to declare an axis; mesh locations (vertices, edges, cells) and index spaces without + geometry subclass `DimensionIndex` directly. Only a declared axis can be staggered, so + `Staggered[Staggered[K]]`, `Staggered[C]` for a mesh location `C` and staggering a local + dimension are errors for the type checkers and at runtime. + + Examples: + >>> class I(CartesianAxisIndex): ... + >>> (I + 1).codomain is I, (I + 0.5).codomain is Staggered[I] + (True, True) + + A mesh location is not an axis: + + >>> class Cell(DimensionIndex): ... + >>> Cell + 1 # doctest: +ELLIPSIS + Traceback (most recent call last): + ... + TypeError: 'Cell' is not a Cartesian axis: ... + >>> Staggered[Cell] # doctest: +ELLIPSIS + Traceback (most recent call last): + ... + TypeError: 'Cell' is not a declared Cartesian axis and cannot be staggered: ... + """ + + __slots__ = () + + _STAGGERED_TAG_RE: Final = re.compile(r"^(?P[^\[\]]+)\[(?P.+)\]$") @@ -783,8 +845,8 @@ def __and__(self, other: Domain) -> Domain: Intersect `Domain`s, missing `Dimension`s are considered infinite. Examples: - >>> class I(DimensionIndex): ... - >>> class J(DimensionIndex): ... + >>> class I(CartesianAxisIndex): ... + >>> class J(CartesianAxisIndex): ... >>> Domain(NamedRange(I, UnitRange(-1, 3))) & Domain(NamedRange(I, UnitRange(1, 6))) Domain(dims=(gt4py.next.common.I[horizontal],), ranges=(UnitRange(1, 3),)) @@ -840,8 +902,8 @@ def slice_at(self) -> utils.IndexerCallable[slice, Domain]: Create a new domain by slicing the domain ranges at the provided relative slices. Examples: - >>> class I(DimensionIndex): ... - >>> class J(DimensionIndex): ... + >>> class I(CartesianAxisIndex): ... + >>> class J(CartesianAxisIndex): ... >>> domain = Domain(NamedRange(I, UnitRange(0, 10)), NamedRange(J, UnitRange(5, 15))) >>> domain.slice_at[2:3, 2:5] Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(2, 3), UnitRange(7, 10))) @@ -932,8 +994,8 @@ def domain(domain_like: DomainLike) -> Domain: Construct `Domain` from `DomainLike` object. Examples: - >>> class I(DimensionIndex): ... - >>> class J(DimensionIndex): ... + >>> class I(CartesianAxisIndex): ... + >>> class J(CartesianAxisIndex): ... >>> domain(((I, (2, 4)), (J, (3, 5)))) Domain(dims=(gt4py.next.common.I[horizontal], gt4py.next.common.J[horizontal]), ranges=(UnitRange(2, 4), UnitRange(3, 5))) @@ -1642,9 +1704,9 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: Examples: >>> from gt4py.next.common import Dimension - >>> class I(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... - >>> class J(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... - >>> class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... + >>> 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): ... >>> promote_dims([J, K], [I, K]) == [I, J, K] @@ -1753,6 +1815,11 @@ def __getitem__(cls, base: Dimension) -> Dimension: raise TypeError( f"'{base.__qualname__}' is already staggered; a dimension cannot be staggered twice." ) + if not issubclass(base, CartesianAxisIndex): + raise TypeError( + f"'{base.__qualname__}' is not a declared Cartesian axis and cannot be staggered:" + " only a dimension declared as 'CartesianAxisIndex' has a staggered partner." + ) if (staggered := _STAGGERED_CACHE.get(base)) is None: staggered = cast( Dimension, @@ -1781,14 +1848,16 @@ def __getitem__(cls, base: Dimension) -> Dimension: if TYPE_CHECKING: # Checkers see an ordinary generic dimension, so `Staggered[K]` works in an annotation and # inside `Field[Dims[Staggered[K]], ...]`. The runtime form below builds a real, interned - # class so that `issubclass` and eve's `type[...]` validation work. Verified clean under - # `mypy --strict` and pyright. - class Staggered[D: DimensionIndex](DimensionIndex): - base: ClassVar[Dimension] + # class so that `issubclass` and eve's `type[...]` validation work. The bound names the + # *declared* level, and `Staggered[K]` is only an `AnyCartesianAxisIndex`, so a doubly + # staggered dimension, a staggered mesh location and a staggered local dimension are all + # `[type-var]` errors. Verified under `mypy --strict` and pyright. + class Staggered[D: CartesianAxisIndex](AnyCartesianAxisIndex): + base: ClassVar[type[CartesianAxisIndex]] else: - class Staggered(DimensionIndex, metaclass=StaggeredMeta): + class Staggered(AnyCartesianAxisIndex, metaclass=StaggeredMeta): """ A dimension sitting at the half-integer positions of a base dimension (ADR 0026). diff --git a/src/gt4py/next/constructors.py b/src/gt4py/next/constructors.py index c07a085fbb..2bef28dce0 100644 --- a/src/gt4py/next/constructors.py +++ b/src/gt4py/next/constructors.py @@ -431,7 +431,7 @@ def empty( Initialize a field in one dimension with a backend and a range domain: >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> a = gtx.empty({IDim: range(3, 10)}, allocator=gtx.itir_python) >>> a.shape (7,) @@ -440,7 +440,7 @@ def empty( >>> import numpy as np >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> a = gtx.empty({IDim: range(3, 10)}, allocator=np) >>> a.shape (7,) @@ -448,7 +448,7 @@ def empty( Initialize with a device and an integer domain. It works like a shape with named dimensions: >>> from gt4py._core import definitions as core_defs - >>> class JDim(gtx.DimensionIndex): ... + >>> class JDim(gtx.CartesianAxisIndex): ... >>> b = gtx.empty( ... {IDim: 3, JDim: 3}, int, device=core_defs.Device(core_defs.DeviceType.CPU, 0) ... ) @@ -476,7 +476,7 @@ def zeros( Examples: >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> gtx.zeros({IDim: range(3, 10)}, allocator=gtx.itir_python).ndarray array([0., 0., 0., 0., 0., 0., 0.]) """ @@ -501,7 +501,7 @@ def ones( Examples: >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> gtx.ones({IDim: range(3, 10)}, allocator=gtx.itir_python).ndarray array([1., 1., 1., 1., 1., 1., 1.]) """ @@ -532,7 +532,7 @@ def full( Examples: >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> gtx.full({IDim: 3}, 5, allocator=gtx.itir_python).ndarray array([5, 5, 5]) """ @@ -577,7 +577,7 @@ def as_field( Examples: >>> import numpy as np >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> xdata = np.array([1, 2, 3]) Automatic domain from just dimensions: diff --git a/src/gt4py/next/embedded/common.py b/src/gt4py/next/embedded/common.py index 3374d30e36..53c6600470 100644 --- a/src/gt4py/next/embedded/common.py +++ b/src/gt4py/next/embedded/common.py @@ -103,7 +103,7 @@ def domain_intersection(*domains: common.Domain) -> common.Domain: Return the intersection of the given domains. Example: - >>> class I(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... >>> domain_intersection( ... common.domain({I: (0, 5)}), common.domain({I: (1, 3)}) ... ) # doctest: +ELLIPSIS @@ -120,8 +120,8 @@ def restrict_to_intersection( Return the with each other intersected domains, ignoring 'ignore_dims' dimensions for the intersection. Example: - >>> class I(common.DimensionIndex): ... - >>> class J(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... + >>> class J(common.CartesianAxisIndex): ... >>> res = restrict_to_intersection( ... common.domain({I: (0, 5), J: (1, 2)}), ... common.domain({I: (1, 3), J: (0, 3)}), diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index 77d87409d4..1428e664d1 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -844,7 +844,7 @@ def scan_operator( >>> import gt4py.next as gtx >>> from gt4py.next.iterator import embedded >>> embedded._column_range = 1 # implementation detail - >>> class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + >>> class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... >>> inp = gtx.as_field([KDim], np.ones((10,))) >>> out = gtx.as_field([KDim], np.zeros((10,))) >>> @gtx.scan_operator(axis=KDim, forward=True, init=0.0) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 3d09d7ecb4..9e8f1fa0f1 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -41,7 +41,7 @@ def with_altered_scalar_kind( >>> print(with_altered_scalar_kind(scalar_t, ts.ScalarKind.BOOL)) bool - >>> class I(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... >>> field_t = ts.FieldType(dims=[I], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) >>> print(with_altered_scalar_kind(field_t, ts.ScalarKind.FLOAT32)) Field[[I], float32] @@ -173,8 +173,8 @@ class FieldOperatorTypeDeduction(traits.VisitorWithSymbolTableTrait, NodeTransla >>> from gt4py.next import Field >>> from gt4py.next.ffront.source_utils import SourceDefinition, get_closure_vars_from_function >>> from gt4py.next.ffront.func_to_foast import FieldOperatorParser - >>> from gt4py.next.common import DimensionIndex - >>> class IDim(DimensionIndex): ... + >>> from gt4py.next.common import CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> def example(a: "Field[[IDim], float]", b: "Field[[IDim], float]"): ... return a + b @@ -699,6 +699,18 @@ def _deduce_binop_type( and type_info.is_arithmetic(right.type) ): # e.g. `IDim+1` or `IDim+0.5` + if not issubclass(left.type.dim, common.AnyCartesianAxisIndex): + raise errors.DSLError( + left.location, + f"'{left.type.dim.__qualname__}' is not a Cartesian axis, so it has no index " + f"arithmetic and '{node.op}' cannot shift along it.", + hints=[ + ( + "Declare a Cartesian axis as 'class IDim(gtx.CartesianAxisIndex): ...';" + " shift along a mesh location with a connectivity instead." + ) + ], + ) if not isinstance(right, foast.Constant): raise errors.DSLError( right.location, diff --git a/src/gt4py/next/ffront/foast_pretty_printer.py b/src/gt4py/next/ffront/foast_pretty_printer.py index 9e2745f0ac..06b4d6b228 100644 --- a/src/gt4py/next/ffront/foast_pretty_printer.py +++ b/src/gt4py/next/ffront/foast_pretty_printer.py @@ -238,8 +238,8 @@ def pretty_format(node: foast.LocatedNode) -> str: Pretty print (to string) an `foast.LocatedNode`. >>> from gt4py.next import Field, Dimension, field_operator, float64 - >>> from gt4py.next.common import DimensionIndex - >>> class IDim(DimensionIndex): ... + >>> from gt4py.next.common import CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> @field_operator ... def field_op(a: Field[[IDim], float64]) -> Field[[IDim], float64]: ... return a + 1.0 diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 929650cf02..f4c8a4fb10 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -80,8 +80,8 @@ class FieldOperatorLowering(eve.PreserveLocationVisitor, eve.NodeTranslator): >>> from gt4py.next.ffront.func_to_foast import FieldOperatorParser >>> from gt4py.next import Field, Dimension, float64 >>> - >>> from gt4py.next.common import DimensionIndex - >>> class IDim(DimensionIndex): ... + >>> from gt4py.next.common import CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> def fieldop(inp: Field[[IDim], "float64"]): ... return inp >>> diff --git a/src/gt4py/next/ffront/foast_to_past.py b/src/gt4py/next/ffront/foast_to_past.py index 64a10a85e8..c18d06cc98 100644 --- a/src/gt4py/next/ffront/foast_to_past.py +++ b/src/gt4py/next/ffront/foast_to_past.py @@ -63,7 +63,7 @@ class OperatorToProgram(workflow.Workflow[ConcreteFOASTOperatorDef, ConcretePAST Example: >>> from gt4py import next as gtx >>> from gt4py.next.otf import arguments, workflow - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> @gtx.field_operator ... def copy(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: diff --git a/src/gt4py/next/ffront/func_to_foast.py b/src/gt4py/next/ffront/func_to_foast.py index 8cc9618ab1..66ad238639 100644 --- a/src/gt4py/next/ffront/func_to_foast.py +++ b/src/gt4py/next/ffront/func_to_foast.py @@ -53,7 +53,7 @@ def func_to_foast(inp: DSLFieldOperatorDef) -> FOASTOperatorDef: Examples: >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> const = gtx.float32(2.0) >>> def dsl_operator(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: @@ -140,8 +140,8 @@ class FieldOperatorParser(DialectParser[foast.FunctionDefinition]): >>> from gt4py.next import Field, Dimension >>> float64 = float - >>> from gt4py.next.common import DimensionIndex - >>> class IDim(DimensionIndex): ... + >>> from gt4py.next.common import CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> def field_op(inp: Field[[IDim], float64]): ... return inp >>> foast_tree = FieldOperatorParser.apply_to_function(field_op) diff --git a/src/gt4py/next/ffront/func_to_past.py b/src/gt4py/next/ffront/func_to_past.py index e4722aa191..77cce10fa3 100644 --- a/src/gt4py/next/ffront/func_to_past.py +++ b/src/gt4py/next/ffront/func_to_past.py @@ -44,7 +44,7 @@ def func_to_past(inp: DSLProgramDef) -> PASTProgramDef: Examples: >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> @gtx.field_operator ... def copy(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index bdedbc3da9..32b9f9dfac 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -42,7 +42,7 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> stages.CompilableProgramDef: Example: >>> from gt4py import next as gtx >>> from gt4py.next.otf import arguments, workflow - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> @gtx.field_operator ... def copy(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: @@ -247,8 +247,8 @@ class ProgramLowering( >>> from gt4py.next import Dimension, Field >>> >>> float64 = float - >>> from gt4py.next.common import DimensionIndex - >>> class IDim(DimensionIndex): ... + >>> from gt4py.next.common import CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> >>> def fieldop(inp: Field[[IDim], "float64"]) -> Field[[IDim], "float64"]: ... >>> def program(inp: Field[[IDim], "float64"], out: Field[[IDim], "float64"]): diff --git a/src/gt4py/next/ffront/type_info.py b/src/gt4py/next/ffront/type_info.py index c99a5fc3e0..fd53a110f2 100644 --- a/src/gt4py/next/ffront/type_info.py +++ b/src/gt4py/next/ffront/type_info.py @@ -188,7 +188,7 @@ def _scan_param_promotion( Example: -------- - >>> class I(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... >>> _scan_param_promotion( ... ts.ScalarType(kind=ts.ScalarKind.INT64), ... ts.FieldType(dims=[I], dtype=ts.ScalarType(kind=ts.ScalarKind.FLOAT64)), diff --git a/src/gt4py/next/field_utils.py b/src/gt4py/next/field_utils.py index 4e2c6add71..3026955162 100644 --- a/src/gt4py/next/field_utils.py +++ b/src/gt4py/next/field_utils.py @@ -36,12 +36,12 @@ def field_from_typespec( The tuple structure and dtype is taken from a type_specifications.DataType, which is either ScalarType or a CollectionTypeSpec of ScalarType (possibly nested). - >>> class I(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... >>> field_from_typespec( ... ts.ScalarType(kind=ts.ScalarKind.INT32), common.domain({I: 1}), np ... ) # doctest: +ELLIPSIS NumPyArrayField(... dtype=int32...) - >>> class I(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... >>> field_from_typespec( ... ts.TupleType( ... types=[ diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 895215d0d7..7b37e35684 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -723,7 +723,7 @@ def _get_axes( In case all arguments are zero-dimensional return an empty sequence. >>> from gt4py import next as gtx - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> i_field: LocatedField = _wrap_field( ... gtx.empty({IDim: range(3, 10)}, allocator=gtx.itir_python) ... ) @@ -731,7 +731,7 @@ def _get_axes( >>> _get_axes((i_field, i_field)) (gt4py.next.iterator.embedded.IDim[horizontal],) - >>> class JDim(gtx.DimensionIndex): ... + >>> class JDim(gtx.CartesianAxisIndex): ... >>> j_field: LocatedField = _wrap_field( ... gtx.empty({JDim: range(3, 10)}, allocator=gtx.itir_python) ... ) diff --git a/src/gt4py/next/iterator/ir_utils/ir_makers.py b/src/gt4py/next/iterator/ir_utils/ir_makers.py index 053e7b78c8..03b7d99e71 100644 --- a/src/gt4py/next/iterator/ir_utils/ir_makers.py +++ b/src/gt4py/next/iterator/ir_utils/ir_makers.py @@ -454,8 +454,8 @@ def domain( ranges_or_domain: dict[common.Dimension, tuple[itir.Expr, itir.Expr]] | common.Domain, ) -> itir.FunCall: """ - >>> class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... - >>> class JDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + >>> class IDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... + >>> class JDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... >>> str(domain(common.GridType.CARTESIAN, {IDim: (0, 10), JDim: (0, 20)})) 'c⟨ gt4py.next.iterator.ir_utils.ir_makers.IDimₕ: [0, 10[, gt4py.next.iterator.ir_utils.ir_makers.JDimₕ: [0, 20[ ⟩' >>> str(domain(common.GridType.UNSTRUCTURED, {IDim: (0, 10), JDim: (0, 20)})) @@ -592,7 +592,7 @@ def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: Examples -------- - >>> class IDim(common.DimensionIndex): ... + >>> class IDim(common.CartesianAxisIndex): ... >>> str(broadcast("a", (IDim,))) 'broadcast(a, {gt4py.next.iterator.ir_utils.ir_makers.IDimₕ})' """ diff --git a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py index de749a356b..3e71508e25 100644 --- a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py +++ b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py @@ -251,7 +251,7 @@ class FuseAsFieldOp( >>> from gt4py import next as gtx >>> from gt4py.next import utils >>> from gt4py.next.iterator.ir_utils import ir_makers as im - >>> class IDim(gtx.DimensionIndex): ... + >>> class IDim(gtx.CartesianAxisIndex): ... >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: >>> import sys diff --git a/src/gt4py/next/iterator/transforms/inline_fundefs.py b/src/gt4py/next/iterator/transforms/inline_fundefs.py index 1158f7ed8e..52c4a8f67e 100644 --- a/src/gt4py/next/iterator/transforms/inline_fundefs.py +++ b/src/gt4py/next/iterator/transforms/inline_fundefs.py @@ -44,7 +44,7 @@ def prune_unreferenced_fundefs(program: itir.Program) -> itir.Program: ... params=[im.sym("a")], ... expr=im.deref("a"), ... ) - >>> class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + >>> class IDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... >>> program = itir.Program( ... id="testee", ... function_definitions=[fun1, fun2], 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 e0db04088a..b9ade1d636 100644 --- a/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py +++ b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py @@ -84,7 +84,7 @@ class _PruneEmptyConcatWhere(PreserveLocationVisitor, NodeTranslator): `gt4py.next.iterator.transforms.concat_where.expand_tuple_args` before to prune them. >>> from gt4py.next import common - >>> class IDim(common.DimensionIndex): ... + >>> class IDim(common.CartesianAxisIndex): ... >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: >>> import sys diff --git a/src/gt4py/next/iterator/transforms/remove_broadcast.py b/src/gt4py/next/iterator/transforms/remove_broadcast.py index f3dcb8a018..7388cf14e2 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -24,11 +24,11 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): Example: >>> from gt4py.next import Dimension, common - >>> from gt4py.next.common import DimensionIndex - >>> class IDim(DimensionIndex): ... + >>> from gt4py.next.common import CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: - >>> class JDim(DimensionIndex): ... + >>> class JDim(CartesianAxisIndex): ... >>> import sys >>> sys.modules[__name__].IDim, sys.modules[__name__].JDim = IDim, JDim >>> domain = im.domain(common.GridType.CARTESIAN, {IDim: (0, 10), JDim: (0, 10)}) diff --git a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py index e7626bac47..95c42360a0 100644 --- a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py +++ b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py @@ -54,7 +54,7 @@ class ReplaceGetDomainRangeWithConstants(PreserveLocationVisitor, NodeTranslator Example: >>> from gt4py import next as gtx - >>> class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + >>> class KDim(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... >>> class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... >>> import sys # register the dimensions where their tags point, as a module would >>> sys.modules[__name__].KDim = KDim diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 098c403755..278683429e 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -479,7 +479,7 @@ def _resolve_dimensions( >>> class Edge(common.DimensionIndex): ... >>> class Vertex(common.DimensionIndex): ... >>> class Cell(common.DimensionIndex): ... - >>> class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... + >>> class K(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... >>> class V2E(common.DimensionIndex): ... >>> class C2V(common.DimensionIndex): ... >>> input_dims = [Edge, K] @@ -508,9 +508,9 @@ def _resolve_dimensions( >>> _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]] >>> from gt4py.next.iterator.ir_utils import ir_makers as im - >>> class IDim(common.DimensionIndex): ... + >>> class IDim(common.CartesianAxisIndex): ... >>> IHalfDim = common.flip_staggered(IDim) - >>> class JDim(common.DimensionIndex): ... + >>> class JDim(common.CartesianAxisIndex): ... >>> # IR passes rebuild a dimension from its tag by importing it (ADR 0029), and a >>> # class declared in a doctest is not an attribute of the real module: >>> import sys diff --git a/src/gt4py/next/type_system/mypy_plugin.py b/src/gt4py/next/type_system/mypy_plugin.py index 35b90168e0..9487baf14b 100644 --- a/src/gt4py/next/type_system/mypy_plugin.py +++ b/src/gt4py/next/type_system/mypy_plugin.py @@ -31,7 +31,7 @@ The documentation on mypy plugins is at https://mypy.readthedocs.io/en/latest/extending_mypy.html Dimensions no longer need plugin support: a concrete dimension is a class -('class IDim(gtx.DimensionIndex): ...'), which is a valid annotation for any type checker. See ADR +('class IDim(gtx.CartesianAxisIndex): ...'), which is a valid annotation for any type checker. See ADR 0029. Only the mixed-precision hooks below remain; this plugin is scheduled for removal once dtype-generic fields land. """ diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index b27625c4f5..453d5f8dad 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -102,7 +102,7 @@ def primitive_constituents( Return the primitive types contained in a composite type. >>> from gt4py.next import common - >>> class I(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... >>> int_type = ts.ScalarType(kind=ts.ScalarKind.INT64) >>> field_type = ts.FieldType(dims=[I], dtype=int_type) @@ -391,8 +391,8 @@ def extract_dims(symbol_type: ts.TypeSpec) -> list[common.Dimension]: Examples: >>> extract_dims(ts.ScalarType(kind=ts.ScalarKind.INT64, shape=[3, 4])) [] - >>> class I(common.DimensionIndex): ... - >>> class J(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... + >>> class J(common.CartesianAxisIndex): ... >>> extract_dims(ts.FieldType(dims=[I, J], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64))) [gt4py.next.type_system.type_info.I[horizontal], gt4py.next.type_system.type_info.J[horizontal]] """ @@ -442,7 +442,7 @@ def is_compatible_type(type_a: ts.TypeSpec, type_b: ts.TypeSpec) -> bool: Beside that this function simply checks for equality of types. >>> bool_type = ts.ScalarType(kind=ts.ScalarKind.BOOL) - >>> class IDim(common.DimensionIndex): ... + >>> class IDim(common.CartesianAxisIndex): ... >>> type_on_i_of_i_it = it_ts.IteratorType( ... position_dims=[IDim], defined_dims=[IDim], element_type=bool_type ... ) @@ -452,7 +452,7 @@ def is_compatible_type(type_a: ts.TypeSpec, type_b: ts.TypeSpec) -> bool: >>> is_compatible_type(type_on_i_of_i_it, type_on_undefined_of_i_it) True - >>> class JDim(common.DimensionIndex): ... + >>> class JDim(common.CartesianAxisIndex): ... >>> type_on_j_of_j_it = it_ts.IteratorType( ... position_dims=[JDim], defined_dims=[JDim], element_type=bool_type ... ) @@ -570,11 +570,11 @@ def promote( `offset_type` of `None` (a list from `make_const_list`) is compatible with any other. >>> dtype = ts.ScalarType(kind=ts.ScalarKind.INT64) - >>> class I(common.DimensionIndex): ... + >>> class I(common.CartesianAxisIndex): ... - >>> class J(common.DimensionIndex): ... + >>> class J(common.CartesianAxisIndex): ... - >>> class K(common.DimensionIndex): ... + >>> class K(common.CartesianAxisIndex): ... >>> promoted: ts.FieldType = promote( ... ts.FieldType(dims=[I, J], dtype=dtype), ts.FieldType(dims=[I, J, K], dtype=dtype), dtype ... ) diff --git a/tests/next_tests/artifacts/custom_named_collections.py b/tests/next_tests/artifacts/custom_named_collections.py index 0007cc12f3..aea5aa7bdc 100644 --- a/tests/next_tests/artifacts/custom_named_collections.py +++ b/tests/next_tests/artifacts/custom_named_collections.py @@ -18,6 +18,7 @@ from gt4py.next import ( common, Dimension, + CartesianAxisIndex, DimensionIndex, Field, float32, @@ -28,7 +29,7 @@ from gt4py.next.type_system import type_specifications as ts -class TDim(DimensionIndex): ... +class TDim(CartesianAxisIndex): ... class SingleElementNamedTupleNamedCollection(NamedTuple): diff --git a/tests/next_tests/benchmarks/benchmark_program_call.py b/tests/next_tests/benchmarks/benchmark_program_call.py index 9f5f90b992..45e0c8bed6 100644 --- a/tests/next_tests/benchmarks/benchmark_program_call.py +++ b/tests/next_tests/benchmarks/benchmark_program_call.py @@ -44,7 +44,7 @@ class Cell(gtx.DimensionIndex): ... -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... @pytest.mark.parametrize("backend", BACKENDS, ids=lambda b: b.name) diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index adc52ee8e2..4b73f163a9 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -160,19 +160,19 @@ def debug_itir(tree): DType = TypeVar("DType") -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... IHalfDim = common.flip_staggered(IDim) -class JDim(gtx.DimensionIndex): ... +class JDim(gtx.CartesianAxisIndex): ... JHalfDim = common.flip_staggered(JDim) -class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... +class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... KHalfDim = common.flip_staggered(KDim) diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py index fa65a38b40..e298515208 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 @@ -30,7 +30,7 @@ from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... I_SIZE = 8 diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py index c5cbd43653..bd1a0bbaa0 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_foast_pretty_printer.py @@ -13,6 +13,7 @@ from gt4py.next import ( Dimension, + CartesianAxisIndex, DimensionIndex, DimensionKind, Field, @@ -26,10 +27,10 @@ from gt4py.next.ffront.func_to_foast import FieldOperatorParser -class I(DimensionIndex): ... +class I(CartesianAxisIndex): ... -class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... +class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... @pytest.mark.parametrize( 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 3700eac616..ce16667142 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 @@ -26,7 +26,7 @@ BACKENDS = [None, gtfn_cpu] -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... @gtx.field_operator diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py index b16461447a..bd6d613efc 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 @@ -73,7 +73,7 @@ def _listify(val): return res -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... def field_maker(*arrays): diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py index ed4039567d..e3a36c77b3 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_conditional.py @@ -16,7 +16,7 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... @fundef diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py index ed1ed9d6a8..1078764bc2 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_implicit_fencil.py @@ -16,7 +16,7 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -class I(gtx.DimensionIndex): ... +class I(gtx.CartesianAxisIndex): ... _isize = 10 diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py index 7220830aab..74dd1dda21 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_program.py @@ -25,7 +25,7 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -class I(gtx.DimensionIndex): ... +class I(gtx.CartesianAxisIndex): ... Ioff = gtx.CartesianConnectivity(I) diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py index 79fcaf1a83..f1cd666b83 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_tuple.py @@ -16,13 +16,13 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... -class JDim(gtx.DimensionIndex): ... +class JDim(gtx.CartesianAxisIndex): ... -class KDim(gtx.DimensionIndex): ... +class KDim(gtx.CartesianAxisIndex): ... # semantics of stencil return that is called from the fencil (after `:` the structure of the output) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py index b370bf7810..485d8fca6b 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_anton_toy.py @@ -24,13 +24,13 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... -class JDim(gtx.DimensionIndex): ... +class JDim(gtx.CartesianAxisIndex): ... -class KDim(gtx.DimensionIndex): ... +class KDim(gtx.CartesianAxisIndex): ... # cross-reference why new type inference does not support this diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py index 746e9f83ed..f6119b3e9f 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_if_stmt.py @@ -22,7 +22,7 @@ def multiply(alpha, inp): return deref(alpha) * deref(inp) -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... @pytest.mark.uses_ir_if_stmts diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py index 6db3cafd2d..2e50346fdc 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_temporaries.py @@ -24,10 +24,10 @@ from next_tests.unit_tests.conftest import program_processor_no_transforms, run_processor -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... -class JDim(gtx.DimensionIndex): ... +class JDim(gtx.CartesianAxisIndex): ... i = gtx.CartesianConnectivity(IDim) diff --git a/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py b/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py index b5588ff7ed..7e34c727aa 100644 --- a/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py +++ b/tests/next_tests/regression_tests/embedded_tests/test_domain_pickle.py @@ -11,10 +11,10 @@ from gt4py.next import common -class I(common.DimensionIndex): ... +class I(common.CartesianAxisIndex): ... -class J(common.DimensionIndex): ... +class J(common.CartesianAxisIndex): ... def test_domain_pickle_after_slice(): diff --git a/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py b/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py index f5a3212889..83d4fd95f6 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_basic_program.py @@ -11,7 +11,7 @@ import gt4py.next as gtx -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... @gtx.field_operator diff --git a/tests/next_tests/unit_tests/embedded_tests/test_common.py b/tests/next_tests/unit_tests/embedded_tests/test_common.py index 2cf2cf3b47..6a7cbf85de 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_common.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_common.py @@ -37,13 +37,13 @@ def test_slice_range(rng, slce, expected): assert result == expected -class I(common.DimensionIndex): ... +class I(common.CartesianAxisIndex): ... -class J(common.DimensionIndex): ... +class J(common.CartesianAxisIndex): ... -class K(common.DimensionIndex): ... +class K(common.CartesianAxisIndex): ... @pytest.mark.parametrize( diff --git a/tests/next_tests/unit_tests/embedded_tests/test_context.py b/tests/next_tests/unit_tests/embedded_tests/test_context.py index 165ec82b2e..734ebcfd8c 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_context.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_context.py @@ -13,10 +13,10 @@ from gt4py.next.errors import exceptions -class IDim(common.DimensionIndex): ... +class IDim(common.CartesianAxisIndex): ... -class NewDim(common.DimensionIndex): ... +class NewDim(common.CartesianAxisIndex): ... def test_getters(): 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 c7b184dcff..6a47073c8a 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 @@ -18,6 +18,7 @@ from gt4py.next import common, constructors from gt4py.next.common import ( Dimension, + CartesianAxisIndex, DimensionIndex, DimensionKind, Domain, @@ -34,13 +35,13 @@ from next_tests.integration_tests.feature_tests.math_builtin_test_data import math_builtin_test_data -class I(DimensionIndex): ... +class I(CartesianAxisIndex): ... -class J(DimensionIndex): ... +class J(CartesianAxisIndex): ... -class I_half(DimensionIndex): ... +class I_half(CartesianAxisIndex): ... class V(DimensionIndex): ... @@ -58,7 +59,7 @@ class V2V(DimensionIndex, kind=DimensionKind.LOCAL): ... class C(DimensionIndex): ... -class K(DimensionIndex): ... +class K(CartesianAxisIndex): ... class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... @@ -70,10 +71,10 @@ class A(DimensionIndex): ... class B(DimensionIndex): ... -class X(DimensionIndex): ... +class X(CartesianAxisIndex): ... -class Y(DimensionIndex): ... +class Y(CartesianAxisIndex): ... class L(DimensionIndex, kind=DimensionKind.LOCAL): ... @@ -103,13 +104,13 @@ class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... class C2V(DimensionIndex): ... -class D0(DimensionIndex): ... +class D0(CartesianAxisIndex): ... -class D1(DimensionIndex): ... +class D1(CartesianAxisIndex): ... -class D2(DimensionIndex): ... +class D2(CartesianAxisIndex): ... @pytest.fixture( diff --git a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py index 25284281ef..2817dd37dc 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 @@ -12,10 +12,10 @@ from gt4py.next.ffront.transform_utils import _deduce_grid_type -class HDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... +class HDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.HORIZONTAL): ... -class VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... +class VDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... class Dim(gtx.DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py index 337c7ce67f..3a272b156e 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py @@ -27,7 +27,7 @@ from gt4py.next.ffront.func_to_foast import FieldOperatorParser -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... IOff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py b/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py index 9018742fe0..c00240655f 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_fbuiltins.py @@ -21,7 +21,7 @@ _SAFE_INPUT = {"arccosh": 2.0} -class IDim(common.DimensionIndex): ... +class IDim(common.CartesianAxisIndex): ... @dataclasses.dataclass 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 0e2cf62732..2b8a8a4c4f 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 @@ -55,7 +55,7 @@ class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -class TDim(gtx.DimensionIndex): ... +class TDim(gtx.CartesianAxisIndex): ... TOff = gtx.FieldOffset("TDim", source=TDim, target=(TDim,)) @@ -69,7 +69,7 @@ class RenamedV2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) -class UDim(gtx.DimensionIndex): ... +class UDim(gtx.CartesianAxisIndex): ... def test_return(): diff --git a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py index 7c3786c28f..3291668c2f 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast.py @@ -54,10 +54,10 @@ from gt4py.next.type_system import type_specifications as ts -class ADim(gtx.DimensionIndex): ... +class ADim(gtx.CartesianAxisIndex): ... -class BDim(gtx.DimensionIndex): ... +class BDim(gtx.CartesianAxisIndex): ... DEREF = itir.SymRef(id=itb.deref.fun.__name__) @@ -78,7 +78,7 @@ class BDim(gtx.DimensionIndex): ... LIFT = itir.SymRef(id=itb.lift.fun.__name__) -class TDim(gtx.DimensionIndex): ... +class TDim(gtx.CartesianAxisIndex): ... # PEP 695 type alias, used to check that aliases are accepted as DSL annotations. diff --git a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py index 330436646b..d1a85cc7c3 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_func_to_foast_error_line_number.py @@ -20,7 +20,7 @@ # NOTE: These tests are sensitive to filename and the line number of the marked statement -class TDim(gtx.DimensionIndex): ... +class TDim(gtx.CartesianAxisIndex): ... def test_invalid_syntax_error_empty_return(): diff --git a/tests/next_tests/unit_tests/ffront_tests/test_stages.py b/tests/next_tests/unit_tests/ffront_tests/test_stages.py index 3509797241..6387edb55b 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_stages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_stages.py @@ -15,7 +15,7 @@ from gt4py.next.type_system import type_specifications as ts -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... def _field_type(kind: ts.ScalarKind) -> ts.FieldType: 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 63fa61893d..ca5cfd4c79 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 @@ -16,6 +16,7 @@ import gt4py.next.ffront.type_specifications from gt4py.next import ( Dimension, + CartesianAxisIndex, DimensionIndex, DimensionKind, Field, @@ -44,22 +45,22 @@ TDim = cnc.TDim -class X(DimensionIndex): ... +class X(CartesianAxisIndex): ... -class Y(DimensionIndex): ... +class Y(CartesianAxisIndex): ... class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... -class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... +class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... -class ADim(DimensionIndex): ... +class ADim(CartesianAxisIndex): ... -class BDim(DimensionIndex): ... +class BDim(CartesianAxisIndex): ... class CDim(DimensionIndex): ... @@ -74,14 +75,14 @@ class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... -class IDim(DimensionIndex): ... +class IDim(CartesianAxisIndex): ... -class JDim(DimensionIndex): ... +class JDim(CartesianAxisIndex): ... # Meaningless dimensions, used for tests. -class SDim(DimensionIndex): ... +class SDim(CartesianAxisIndex): ... def test_unpack_assign(): @@ -584,6 +585,15 @@ def as_offset_cross_dim(a: Field[[IDim], float], b: Field[[IDim], int]): _ = FieldOperatorParser.apply_to_function(as_offset_cross_dim) +@pytest.mark.parametrize("offset", [1, -1, 0.5]) +def test_cartesian_shift_off_an_axis(offset): + def shift_along_mesh_location(a: Field[[Edge], float]): + return a(Edge + offset) + + with pytest.raises(errors.DSLError, match="'Edge' is not a Cartesian axis"): + _ = FieldOperatorParser.apply_to_function(shift_along_mesh_location) + + vpfloat: TypeAlias = float32 wpfloat: TypeAlias = float64 diff --git a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py index 7120882c75..3c7157166d 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 @@ -15,16 +15,16 @@ from gt4py.next import common, constructors -class I(common.DimensionIndex): ... +class I(common.CartesianAxisIndex): ... IHalf = common.flip_staggered(I) -class J(common.DimensionIndex): ... +class J(common.CartesianAxisIndex): ... -class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class K(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... class Vertex(common.DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py index 894a85fa1d..9e7e73862a 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py @@ -17,7 +17,7 @@ from gt4py.next.iterator import embedded -class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class K(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... def test_column_ufunc(): diff --git a/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py b/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py index b5432098bc..4c7b568f1e 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py @@ -12,7 +12,7 @@ from gt4py.next.type_system import type_specifications as ts -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... field_type = ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.INT32)) 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 923bda3ff4..7119c5fcb7 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 @@ -261,10 +261,10 @@ def test_make_tuple(): assert actual == expected -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... -class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... +class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... 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 3cba7b16cd..a742561830 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 @@ -37,7 +37,7 @@ def foo(inp): ) -class I(gtx.DimensionIndex): ... +class I(gtx.CartesianAxisIndex): ... def test_deduce_domain(): diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py index b6a7d05388..7a8570949e 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_canonicalize_domain_args.py @@ -18,7 +18,7 @@ int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) -class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class IDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... field_type = ts.FieldType(dims=[IDim], dtype=int_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py index 97c2f4ce40..c476dc4ee8 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py @@ -19,7 +19,7 @@ int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) -class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class IDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... field_type = ts.FieldType(dims=[IDim], dtype=int_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py index babcafc4a1..3454b36db5 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_transform_to_as_fieldop.py @@ -16,10 +16,10 @@ int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) -class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class IDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... -class JDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class JDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... def test_in_helper(): diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py index e2e73a819f..a0357b9976 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py @@ -20,7 +20,7 @@ ) -class I(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class I(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... @pytest.fixture diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py index f2e4f4bb07..bc1dbc3355 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py @@ -15,7 +15,7 @@ from gt4py.next.iterator.transforms import dead_code_elimination -class TDim(common.DimensionIndex): ... +class TDim(common.CartesianAxisIndex): ... int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py index eafe45b9d3..ff6890f906 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 @@ -30,13 +30,13 @@ float_type = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) -class IDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class IDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... -class JDim(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class JDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... -class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class KDim(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py index d8f3327386..8e9a0be03a 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py @@ -16,7 +16,7 @@ from gt4py.next.type_system import type_specifications as ts -class IDim(common.DimensionIndex): ... +class IDim(common.CartesianAxisIndex): ... T = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py index 5bceba0e53..250d3eabc3 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 @@ -20,10 +20,10 @@ class Neighbor(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... -class IDim(common.DimensionIndex): ... +class IDim(common.CartesianAxisIndex): ... -class JDim(common.DimensionIndex): ... +class JDim(common.CartesianAxisIndex): ... field_type = ts.FieldType(dims=[IDim], dtype=ts.ScalarType(kind=ts.ScalarKind.INT32)) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py index 9adfdaedbd..0c39338b8d 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py @@ -26,13 +26,13 @@ ) -class IDim(common.DimensionIndex): ... +class IDim(common.CartesianAxisIndex): ... -class JDim(common.DimensionIndex): ... +class JDim(common.CartesianAxisIndex): ... -class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class KDim(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... index_type = ts.ScalarType(kind=getattr(ts.ScalarKind, builtins.INTEGER_INDEX_BUILTIN.upper())) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py index 3a7c92b784..f391363e76 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py @@ -14,7 +14,7 @@ from gt4py.next.iterator.ir_utils import ir_makers as im -class TDim(common.DimensionIndex): ... +class TDim(common.CartesianAxisIndex): ... int_type = ts.ScalarType(kind=ts.ScalarKind.INT32) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py index 26cff82c6c..a745e08847 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py @@ -13,7 +13,7 @@ from gt4py.next.type_system import type_specifications as ts -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... def test_prune_casts_simple(): diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py index 3f6cd212ae..e4d1746ccc 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 @@ -28,7 +28,7 @@ class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... -class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class K(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... float64 = ts.ScalarType(kind=ts.ScalarKind.FLOAT64) diff --git a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py index 7160b12de4..4b86bfbe8e 100644 --- a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py +++ b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/conftest.py @@ -19,10 +19,10 @@ from gt4py.next.otf.compilation import cache -class I(gtx.DimensionIndex): ... +class I(gtx.CartesianAxisIndex): ... -class J(gtx.DimensionIndex): ... +class J(gtx.CartesianAxisIndex): ... def make_program_source(name: str) -> artifacts.ProgramSource: diff --git a/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py b/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py index bd1be39266..50f0e31f24 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py +++ b/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py @@ -62,7 +62,7 @@ def test_sanitize_static_args_wrong_type(): static_arg.validate("foo", ts.ScalarType(kind=ts.ScalarKind.INT32)) -class TDim(gtx.DimensionIndex): ... +class TDim(gtx.CartesianAxisIndex): ... @pytest.fixture diff --git a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py index d233f29a8d..2e27d05687 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py +++ b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py @@ -30,7 +30,7 @@ ) -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... @pytest.fixture diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py index b08ff90f4f..d2f2c47c75 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/transformation_tests/test_map_promoter.py @@ -22,10 +22,10 @@ from . import util -class boden(gtx_common.DimensionIndex): ... +class boden(gtx_common.CartesianAxisIndex): ... -class K(gtx_common.DimensionIndex, kind=gtx_common.DimensionKind.VERTICAL): ... +class K(gtx_common.CartesianAxisIndex, kind=gtx_common.DimensionKind.VERTICAL): ... N = 10 diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index a2c4832b78..e98c5a6517 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -20,6 +20,7 @@ import gt4py.next.common as common from gt4py.next.common import ( Dimension, + CartesianAxisIndex, DimensionIndex, DimensionKind, Domain, @@ -33,28 +34,28 @@ ) -class X(DimensionIndex): ... +class X(CartesianAxisIndex): ... -class Y(DimensionIndex): ... +class Y(CartesianAxisIndex): ... -class Z(DimensionIndex): ... +class Z(CartesianAxisIndex): ... class Foo(DimensionIndex): ... -class J(DimensionIndex): ... +class J(CartesianAxisIndex): ... -class K(DimensionIndex): ... +class K(CartesianAxisIndex): ... -class I(common.DimensionIndex): ... +class I(common.CartesianAxisIndex): ... -class I_half(common.DimensionIndex): ... +class I_half(common.CartesianAxisIndex): ... class C2E(DimensionIndex, kind=DimensionKind.LOCAL): ... @@ -75,13 +76,13 @@ class E2C2V(DimensionIndex, kind=DimensionKind.LOCAL): ... class ECDim(DimensionIndex): ... -class IDim(DimensionIndex): ... +class IDim(CartesianAxisIndex): ... -class JDim(DimensionIndex): ... +class JDim(CartesianAxisIndex): ... -class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... +class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... IHalfDim = common.flip_staggered(IDim) @@ -914,6 +915,24 @@ def test_a_dimension_cannot_be_staggered_twice(self): with pytest.raises(TypeError, match="is already staggered"): common.Staggered[KDim][KDim] + @pytest.mark.parametrize( + "dim, match", + [ + (ECDim, "is not a declared Cartesian axis"), # a mesh location + (V2E, "is not a declared Cartesian axis"), # a local dimension + (DimensionIndex, "is not a declared Cartesian axis"), + (common.AnyCartesianAxisIndex, "is not a declared Cartesian axis"), + ], + ) + def test_only_a_declared_axis_can_be_staggered(self, dim, match): + with pytest.raises(TypeError, match=match): + common.Staggered[dim] + + def test_a_staggered_dimension_is_an_axis_but_not_a_declared_one(self): + assert issubclass(common.Staggered[KDim], common.AnyCartesianAxisIndex) + assert not issubclass(common.Staggered[KDim], common.CartesianAxisIndex) + assert not issubclass(common.Staggered[KDim], KDim) + def test_resolve_rejects_a_bracketed_tag_of_another_owner(self): with pytest.raises(ValueError, match="not a parametrized dimension"): common.resolve(f"{KDim.tag}[{KDim.tag}]") @@ -926,3 +945,28 @@ def test_resolve_loaded(): assert common.resolve_loaded("not_imported_anywhere.IDim") is None assert common.resolve_loaded(f"{__name__}.does_not_exist") is None assert common.resolve_loaded(f"{__name__}.test_resolve_loaded") is None # not a dimension + + +class TestCartesianAxisIndex: + def test_axis_levels(self): + assert issubclass(IDim, common.CartesianAxisIndex) + assert issubclass(common.CartesianAxisIndex, common.AnyCartesianAxisIndex) + assert issubclass(common.AnyCartesianAxisIndex, DimensionIndex) + assert not issubclass(ECDim, common.AnyCartesianAxisIndex) + + @pytest.mark.parametrize("dim", [IDim, IHalfDim]) + def test_shift_along_an_axis(self, dim): + assert (dim + 1).codomain is dim + assert (dim - 1).codomain is dim + assert (dim + 0.5).codomain is common.flip_staggered(dim) + + @pytest.mark.parametrize("dim", [ECDim, V2E]) + @pytest.mark.parametrize("op", [lambda d: d + 1, lambda d: d - 1, lambda d: d + 0.5]) + def test_no_index_arithmetic_off_an_axis(self, dim, op): + with pytest.raises(TypeError, match="is not a Cartesian axis"): + op(dim) + + def test_comparisons_stay_on_every_dimension(self): + # `concat_where(EdgeDim < n, ...)` over a mesh location must keep working + assert (ECDim < 5) == Domain(dims=(ECDim,), ranges=(UnitRange(common.Infinity.NEGATIVE, 5),)) + assert (ECDim == 3) == Domain(dims=(ECDim,), ranges=(UnitRange(3, 4),)) diff --git a/tests/next_tests/unit_tests/test_constructors.py b/tests/next_tests/unit_tests/test_constructors.py index 3bfdacbf63..b84d5d303f 100644 --- a/tests/next_tests/unit_tests/test_constructors.py +++ b/tests/next_tests/unit_tests/test_constructors.py @@ -22,13 +22,13 @@ ) -class I(gtx.DimensionIndex): ... +class I(gtx.CartesianAxisIndex): ... -class J(gtx.DimensionIndex): ... +class J(gtx.CartesianAxisIndex): ... -class K(gtx.DimensionIndex): ... +class K(gtx.CartesianAxisIndex): ... sizes = {I: 10, J: 10, K: 10} diff --git a/tests/next_tests/unit_tests/test_custom_layout_allocators.py b/tests/next_tests/unit_tests/test_custom_layout_allocators.py index f5257506fb..47c657c3d3 100644 --- a/tests/next_tests/unit_tests/test_custom_layout_allocators.py +++ b/tests/next_tests/unit_tests/test_custom_layout_allocators.py @@ -17,28 +17,28 @@ import gt4py.storage.allocators as core_allocators -class D0(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class D0(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... -class D1(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class D1(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... -class D2(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... +class D2(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... -class D0_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class D0_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... -class D2_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class D2_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... -class D1_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class D1_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... class DummyAllocator(next_allocators.FieldBufferAllocatorProtocol): @@ -154,7 +154,7 @@ class Cell(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... +class K(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... class TestBaseFieldBufferAllocatorAlignedIndex: diff --git a/tests/next_tests/unit_tests/test_field_utils.py b/tests/next_tests/unit_tests/test_field_utils.py index a26ffa2992..54e73c6a51 100644 --- a/tests/next_tests/unit_tests/test_field_utils.py +++ b/tests/next_tests/unit_tests/test_field_utils.py @@ -12,7 +12,7 @@ from gt4py.next import common, constructors, field_utils -class X(common.DimensionIndex): ... +class X(common.CartesianAxisIndex): ... @pytest.mark.parametrize( diff --git a/tests/next_tests/unit_tests/test_utils.py b/tests/next_tests/unit_tests/test_utils.py index 6ad27a6f59..aad5186a85 100644 --- a/tests/next_tests/unit_tests/test_utils.py +++ b/tests/next_tests/unit_tests/test_utils.py @@ -18,10 +18,10 @@ from eve_tests import definitions -class I(common.DimensionIndex): ... +class I(common.CartesianAxisIndex): ... -class J(common.DimensionIndex): ... +class J(common.CartesianAxisIndex): ... @dataclasses.dataclass diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py index 6ee13a358c..2984713b76 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 @@ -12,6 +12,7 @@ from gt4py.next import ( Dimension, + CartesianAxisIndex, DimensionIndex, DimensionKind, ) @@ -20,13 +21,13 @@ from gt4py.next.iterator.type_system import type_specifications as ts_it -class IDim(DimensionIndex): ... +class IDim(CartesianAxisIndex): ... -class JDim(DimensionIndex): ... +class JDim(CartesianAxisIndex): ... -class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... +class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... @@ -35,7 +36,7 @@ class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... -class TDim(DimensionIndex): ... +class TDim(CartesianAxisIndex): ... def type_info_cases() -> list[tuple[Optional[ts.TypeSpec], dict]]: diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py index 7ba7d09b5b..adbde674ec 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_translation.py @@ -30,10 +30,10 @@ def dtype(self) -> np.dtype: return np.dtype(np.int32) -class IDim(gtx.DimensionIndex): ... +class IDim(gtx.CartesianAxisIndex): ... -class JDim(gtx.DimensionIndex): ... +class JDim(gtx.CartesianAxisIndex): ... # -- PEP 695 type aliases -- diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 20ed020333..1506c46a75 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -4,7 +4,7 @@ import typing from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -16,7 +16,7 @@ main: | from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator(grid_type=None, backend=None) def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -26,7 +26,7 @@ main: | from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator(grid_type=None) def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -36,7 +36,7 @@ main: | from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.scan_operator(axis=KDim) def somescanop( @@ -49,7 +49,7 @@ main: | from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.scan_operator(axis=KDim, forward=True, init=0.0) def somescanop( @@ -64,7 +64,7 @@ import typing from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -85,7 +85,7 @@ import typing from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -106,7 +106,7 @@ import typing from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -128,7 +128,7 @@ from gt4py import next as gtx class CellDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... T = typing.TypeVar("T", float, gtx.float32, gtx.float64, bool, gtx.int32, gtx.int64) CellKField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim, KDim], T] @@ -146,7 +146,7 @@ from gt4py import next as gtx class CellDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... T = typing.TypeVar("T", float, gtx.float32, gtx.float64, bool, gtx.int32, gtx.int64) CellKField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim, KDim], T] CKTuple: typing.TypeAlias = tuple[CellKField[T], CellKField[T]] @@ -164,7 +164,7 @@ from gt4py.next import where class CellDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... T = typing.TypeVar("T", float, gtx.float32, gtx.float64, bool, gtx.int32, gtx.int64) CellField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim], T] CellKField: typing.TypeAlias = gtx.Field[gtx.Dims[CellDim, KDim], T] @@ -205,7 +205,7 @@ import typing from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -218,7 +218,7 @@ main: | from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.scan_operator(axis=KDim) def somescanop( @@ -236,7 +236,7 @@ import typing from gt4py import next as gtx - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo(a: gtx.Field[gtx.Dims[KDim], gtx.int32]) -> gtx.Field[gtx.Dims[KDim], gtx.int32]: @@ -261,7 +261,7 @@ from gt4py import next as gtx class CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.HORIZONTAL): ... - class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... @gtx.field_operator def foo( @@ -275,3 +275,30 @@ main: | import xarray a: xarray.NamedArray + + - case: cartesian_axis_levels + main: | + from gt4py import next as gtx + + class K(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... + class C(gtx.DimensionIndex): ... + + def any_dimension(d: gtx.Dimension) -> None: ... + def any_axis(d: type[gtx.AnyCartesianAxisIndex]) -> None: ... + + any_dimension(gtx.Staggered[K]) + any_axis(gtx.Staggered[K]) + ok_dims: gtx.Field[gtx.Dims[C, gtx.Staggered[K]], float] | None = None + ok_shift = K + 1 + ok_half = gtx.Staggered[K] - 0.5 + any_axis(C) + bad_doubly: gtx.Field[gtx.Dims[gtx.Staggered[gtx.Staggered[K]]], float] | None = None + bad_location: gtx.Field[gtx.Dims[gtx.Staggered[C]], float] | None = None + bad_shift = C + 1 + bad_shift_back = C - 1 + out: | + main:14:10: error: Argument 1 to "any_axis" has incompatible type "type[C]"; expected "type[AnyCartesianAxisIndex]" [arg-type] + main:15:46: error: Type argument "Staggered[K]" of "Staggered" must be a subtype of "CartesianAxisIndex" [type-var] + 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] From ea75ee7868196021722432007a52be70487a9d32 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 00:41:56 +0200 Subject: [PATCH 19/20] fix[next]: fingerprint a dimension's kind A dimension's `kind` decides a field's layout order and the scan axis, so a dimension redefined under the same name with another kind must not reuse compiled artifacts. Staggered dimensions follow through their base. --- .../next/0029-Dimensions_As_Nominal_Types.md | 5 +++- src/gt4py/next/fingerprinting.py | 11 +++++++++ tests/next_tests/unit_tests/test_common.py | 24 ++++++++++++++++++- 3 files changed, 38 insertions(+), 2 deletions(-) diff --git a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md index f1ec832a45..3f901f8f59 100644 --- a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md @@ -120,7 +120,10 @@ disappears. 4. **Cache fingerprints depend on module paths.** A dimension is fingerprinted by qualified name, so moving a declaration between modules invalidates compiled - artifacts. This is a consequence for the build cache of ADR 0023, not a + artifacts. Its `kind` is fingerprinted too: it decides a field's layout order and + the scan axis, so a dimension redefined under the same name with another `kind` + (a re-run notebook cell) does not reuse artifacts. A staggered dimension is + fingerprinted through its base, and its base is fixed by interning. This is a consequence for the build cache of ADR 0023, not a reversal of it. The generic `type` deconstructor is correct for the *lenient* fingerprint variant; the STRICT variant rejects a parametrized dimension, which is not importable under its qualified name. diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index 27316eaac5..52a7081b1d 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -335,8 +335,18 @@ def object_deconstruct_fallback(obj: Any) -> Deconstruction: ) +def _dimension_deconstruction(obj: common.DimensionMeta, *, strict: bool = True) -> Deconstruction: + # NOTE: by reference, like any class, plus the declaration-time `kind`. It decides a field's + # layout order (`order_dimensions`) and the scan axis, so a dimension redefined under the same + # name with another `kind` (a re-run notebook cell) must not reuse compiled artifacts. A + # staggered dimension goes through its base (see above), so it inherits this. See ADR 0029. + reference = Deconstruction.from_reference(obj, strict=strict).state + return Deconstruction.from_pieces(obj.kind, state=b"dimension\0" + reference) + + #: Strict deconstructors map used by `strict_fingerprinter` STRICT_DECONSTRUCTORS: Final[dict[type, Deconstructor]] = _COMMON_DECONSTRUCTORS | { + common.DimensionMeta: _dimension_deconstruction, type: EmptyDeconstruction.from_reference, types.FunctionType: EmptyDeconstruction.from_reference, types.BuiltinFunctionType: EmptyDeconstruction.from_reference, @@ -396,6 +406,7 @@ def _lenient_function_deconstruction(func: types.FunctionType) -> Deconstruction #: Tolerant deconstructors map used by `lenient_fingerprinter` LENIENT_DECONSTRUCTORS: Final[dict[type, Deconstructor]] = { + common.DimensionMeta: functools.partial(_dimension_deconstruction, strict=False), types.FunctionType: _lenient_function_deconstruction, types.BuiltinFunctionType: _lenient_reference, types.ModuleType: _lenient_reference, diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index e98c5a6517..c5daa24c9b 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -968,5 +968,27 @@ def test_no_index_arithmetic_off_an_axis(self, dim, op): def test_comparisons_stay_on_every_dimension(self): # `concat_where(EdgeDim < n, ...)` over a mesh location must keep working - assert (ECDim < 5) == Domain(dims=(ECDim,), ranges=(UnitRange(common.Infinity.NEGATIVE, 5),)) + assert (ECDim < 5) == Domain( + dims=(ECDim,), ranges=(UnitRange(common.Infinity.NEGATIVE, 5),) + ) assert (ECDim == 3) == Domain(dims=(ECDim,), ranges=(UnitRange(3, 4),)) + + +def test_a_dimension_fingerprint_includes_its_kind(): + # `kind` sets the layout order, so a dimension redefined under the same name with another kind + # (a re-run notebook cell) must not reuse artifacts; a staggered dimension follows its base. + from gt4py.next import fingerprinting + + def declare(source: str) -> type: + namespace = {"__name__": __name__, "common": common} + exec(source, namespace) + return namespace["Redefined"] + + horizontal = declare("class Redefined(common.CartesianAxisIndex): ...") + vertical = declare( + "class Redefined(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ..." + ) + fingerprint = fingerprinting.lenient_fingerprinter + assert fingerprint(horizontal) != fingerprint(vertical) + assert fingerprint(common.Staggered[horizontal]) != fingerprint(common.Staggered[vertical]) + assert fingerprinting.strict_fingerprinter(KDim) == fingerprinting.strict_fingerprinter(KDim) From a1cf830095afae2da13db41427581901a244c128 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 01:34:32 +0200 Subject: [PATCH 20/20] docs[next]: mdformat the Cartesian axis table of ADR 0029 --- .../ADRs/next/0029-Dimensions_As_Nominal_Types.md | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md index 3f901f8f59..fcb29cbd04 100644 --- a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md @@ -205,11 +205,11 @@ Both levels sit *below* `DimensionIndex`, so `Staggered[K]` stays a `DimensionIn and no annotation or `issubclass` guard in the tree widens. The bound is on the *declared* level, and `Staggered[K]` is only an `AnyCartesianAxisIndex`, so: -| Rejected | Statically (mypy, pyright) | At runtime | -| --- | --- | --- | -| `Staggered[Staggered[K]]` | `[type-var]` | `TypeError` | -| `Staggered[Cell]`, staggering a local dimension | `[type-var]` | `TypeError` | -| `Cell + 1`, `Cell - 1` | `[operator]` | `TypeError`; a `DSLError` in a field operator | +| Rejected | Statically (mypy, pyright) | At runtime | +| ----------------------------------------------- | -------------------------- | --------------------------------------------- | +| `Staggered[Staggered[K]]` | `[type-var]` | `TypeError` | +| `Staggered[Cell]`, staggering a local dimension | `[type-var]` | `TypeError` | +| `Cell + 1`, `Cell - 1` | `[operator]` | `TypeError`; a `DSLError` in a field operator | The last row needs `DimensionMeta.__add__` / `__sub__` declared with the self-type `cls: type[AnyCartesianAxisIndex]`. Both checkers bind it correctly at every call