diff --git a/docs/development/ADRs/next/0026-Staggered_Dimensions.md b/docs/development/ADRs/next/0026-Staggered_Dimensions.md index c8fbd3a91f..7584508c8d 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-10-02 + +> The *encoding* of this record is superseded by +> [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. 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,14 @@ 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 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`. 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 @@ -87,7 +97,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/ADRs/next/0029-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md new file mode 100644 index 0000000000..fcb29cbd04 --- /dev/null +++ b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md @@ -0,0 +1,343 @@ +--- +tags: [] +--- + +# Dimensions as Nominal Types + +- **Status**: proposed +- **Authors**: Enrique González Paredes (@egparedes) +- **Created**: 2026-09-18 +- **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.CartesianAxisIndex): ... + + +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` +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__}"`. A Cartesian axis — a dimension with +index arithmetic and a staggered partner — is a subclass of `CartesianAxisIndex`; +mesh locations subclass `DimensionIndex` directly. + +## 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`. 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 + 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`. + + **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 + 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. 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. + +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").replace("[", "_l").replace("]", "_r") + + + def from_codegen_name(name: str) -> Tag: + 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. 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 + `_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. + `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 + 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. + +### 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 +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. + +### 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. + +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. + +### 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. 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. +- **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 + +- 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). +- 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/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index f19167ef9b..4f1661d745 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) +- [0029 - Dimensions as Nominal Types](0029-Dimensions_As_Nominal_Types.md) ### Frontend and Parsing #frontend diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 829acfeffa..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 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. `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 -CellDim = gtx.Dimension("Cell") -KDim = gtx.Dimension("K") +class CellDim(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 -I = gtx.Dimension("I") -J = gtx.Dimension("J") +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( @@ -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..dc2895dc4c 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.CartesianAxisIndex): ...\n", + "\n", + "\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 dd1a30cc6b..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,8 +51,12 @@ "metadata": {}, "outputs": [], "source": [ - "I = gtx.Dimension(\"I\")\n", - "J = gtx.Dimension(\"J\")\n", + "class I(gtx.CartesianAxisIndex): ...\n", + "\n", + "\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 8fd90daf20..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, 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, @@ -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(CartesianAxisIndex, 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..643950aa9c 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.CartesianAxisIndex, 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..6a44a6eef7 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.CartesianAxisIndex, 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..a85937180f 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.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 a5b8aaf78f..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": [ - "K = gtx.Dimension(\"K\", 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 9a8dfc92b5..fd4a4338c8 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.CartesianAxisIndex): ...\n", + "\n", + "\n", + "class J(gtx.CartesianAxisIndex): ...\n", + "\n", + "\n", + "class K(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ...\n", + "\n", "\n", "domain = gtx.domain({I: nx, J: ny, K: nz})\n", "\n", diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index 3b9f97592f..b8a7bf5143 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -22,19 +22,24 @@ 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, + 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 +120,12 @@ "is_scalar_type", # from common "Dimension", + "DimensionIndex", + "AnyCartesianAxisIndex", + "CartesianAxisIndex", "DimensionKind", + "Staggered", + "resolve", "Dims", "Field", "CartesianConnectivity", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 5bc53f6474..a5486c024f 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -10,10 +10,13 @@ import abc import collections +import copyreg import dataclasses import enum import functools +import importlib import math +import re import sys import types from collections.abc import Callable, Iterable, Mapping, Sequence @@ -60,6 +63,62 @@ 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. + + 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' + """ + # 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: + """ + 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"_([udlr])", lambda m: _CODEGEN_UNESCAPE[m.group(1)], name) + + @enum.unique class DimensionKind(StrEnum): HORIZONTAL = "horizontal" @@ -73,55 +132,115 @@ 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. - def __call__(self, val: int) -> NamedIndex: - return NamedIndex(self, val) + 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 __add__(self, offset: int | float) -> Connectivity: - return connectivity_for_cartesian_shift(self, offset) + def __repr__(cls) -> str: + return f"{cls.tag}[{cls.kind}]" + + 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}]" + + # 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__(self, offset: int | float) -> Connectivity: - return self + (-offset) + def __sub__( # type: ignore[misc] + cls: type[AnyCartesianAxisIndex], # pyright: ignore[reportGeneralTypeIssues] + offset: int | float, + ) -> Connectivity: + return cls + (-offset) - def __gt__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(value + 1, 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 __ge__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(value, Infinity.POSITIVE),)) + def __ge__(cls: Dimension, value: core_defs.IntegralScalar) -> Domain: # type: ignore[misc] + return Domain(dims=(cls,), ranges=(UnitRange(value, Infinity.POSITIVE),)) - def __lt__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(Infinity.NEGATIVE, value),)) + def __lt__(cls: Dimension, value: core_defs.IntegralScalar) -> Domain: # type: ignore[misc] + return Domain(dims=(cls,), ranges=(UnitRange(Infinity.NEGATIVE, value),)) - def __le__(self, value: core_defs.IntegralScalar) -> Domain: - return Domain(dims=(self,), ranges=(UnitRange(Infinity.NEGATIVE, value + 1),)) + 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 supertype `object.__eq__` which returns `bool`. - def __eq__(self, value: Dimension) -> bool: ... + @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 0029 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=(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, " @@ -131,26 +250,238 @@ 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(CartesianAxisIndex): ... + >>> class K(CartesianAxisIndex, 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(CartesianAxisIndex): ... + >>> 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 0029). + 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 0029. +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__ = () - @dataclasses.dataclass(frozen=True) - class _DimA(Dimension): ... - @dataclasses.dataclass(frozen=True) - class _DimB(Dimension): ... +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: - @dataclasses.dataclass(frozen=True) - class _DimC(Dimension): ... + >>> 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: ... + """ - @dataclasses.dataclass(frozen=True) - class _DimD(Dimension): ... + __slots__ = () - @dataclasses.dataclass(frozen=True) - class _AnyDim(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: + """ + 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. + + 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. + + 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. + + Args: + tag: A dimension tag, as produced by `DimensionMeta.tag`. + + Returns: + The dimension class. + + 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"]) + 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): + 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 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." + ) + + +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): @@ -364,20 +695,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 @@ -397,16 +720,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]: @@ -448,7 +771,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}'." ) @@ -503,7 +826,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: @@ -522,16 +845,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(CartesianAxisIndex): ... + >>> class J(CartesianAxisIndex): ... >>> 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( @@ -579,10 +902,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(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=(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: @@ -632,7 +956,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.") @@ -670,20 +994,20 @@ def domain(domain_like: DomainLike) -> Domain: Construct `Domain` from `DomainLike` object. Examples: - >>> I = Dimension("I") - >>> J = Dimension("J") + >>> class I(CartesianAxisIndex): ... + >>> class J(CartesianAxisIndex): ... >>> 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 @@ -739,7 +1063,10 @@ def __gt_domain__(self) -> Domain: @property def __gt_dims__(self) -> tuple[str, ...]: - return tuple(d.value 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 0029). + return tuple(d.__qualname__ for d in self.__gt_domain__.dims) @runtime_checkable @@ -1333,11 +1660,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, ), ) @@ -1367,15 +1700,15 @@ 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 - >>> 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(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] True >>> promote_dims([K, J], [I, K]) @@ -1439,20 +1772,164 @@ 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 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. +_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 0029. + """ + + #: 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 + + @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"{cls.__module__}.Staggered[{cls.base.tag}]" + return super().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." + ) + if not isinstance(base, DimensionMeta): + raise TypeError(f"'Staggered' expects a dimension, got '{base!r}'.") + if is_staggered(base): + 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, + 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.__qualname__}]", + }, + ), + ) + # 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: + # 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. 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(AnyCartesianAxisIndex, 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 ConstList(DimensionIndex, kind=DimensionKind.LOCAL): + """ + The local dimension of a list of one repeated value (`make_const_list`). + + The value is broadcast against the neighbor lists it is combined with, and a materialized + constant list has extent 1 along it. It indexes no table, so it is never in an offset provider. + + Declared here, once: it used to be built independently in `iterator/embedded.py` and in the + DaCe lowering, which only worked while dimensions compared by `(name, kind)`. + """ + + __slots__ = () + + +def _reduce_staggered(cls: StaggeredMeta) -> Any: + """ + Pickle a staggered dimension through its base, falling back to by-reference. + + `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) 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/constructors.py b/src/gt4py/next/constructors.py index e4320f99d3..2bef28dce0 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: @@ -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.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 - >>> IDim = gtx.Dimension("I") + >>> 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 - >>> JDim = gtx.Dimension("J") + >>> class JDim(gtx.CartesianAxisIndex): ... >>> b = gtx.empty( ... {IDim: 3, JDim: 3}, int, device=core_defs.Device(core_defs.DeviceType.CPU, 0) ... ) @@ -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: @@ -476,7 +476,7 @@ def zeros( Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.CartesianAxisIndex): ... >>> gtx.zeros({IDim: range(3, 10)}, allocator=gtx.itir_python).ndarray array([0., 0., 0., 0., 0., 0., 0.]) """ @@ -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: @@ -501,7 +501,7 @@ def ones( Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.CartesianAxisIndex): ... >>> gtx.ones({IDim: range(3, 10)}, allocator=gtx.itir_python).ndarray array([1., 1., 1., 1., 1., 1., 1.]) """ @@ -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: @@ -532,7 +532,7 @@ def full( Examples: >>> from gt4py import next as gtx - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.CartesianAxisIndex): ... >>> gtx.full({IDim: 3}, 5, allocator=gtx.itir_python).ndarray array([5, 5, 5]) """ @@ -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: @@ -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.CartesianAxisIndex): ... >>> xdata = np.array([1, 2, 3]) Automatic domain from just dimensions: @@ -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, @@ -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/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..53c6600470 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( @@ -97,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.CartesianAxisIndex): ... >>> 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())) @@ -114,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.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)}), @@ -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..b8133cbafa 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.__qualname__ for d in offset.target) # for the diagnostic raise ValueError( f"'as_offset' is only supported for Cartesian offsets " f"(single target dimension equal to source dimension); " - f"got source '{offset.source.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/decorator.py b/src/gt4py/next/ffront/decorator.py index bf63aa692a..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 - >>> KDim = gtx.Dimension("K", 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/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 2fd072a483..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, @@ -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..9e8f1fa0f1 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.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] """ @@ -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 CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> def example(a: "Field[[IDim], float]", b: "Field[[IDim], float]"): ... return a + b @@ -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): @@ -697,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, @@ -710,7 +724,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 +1003,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.__qualname__ for d in arg_0.target) # for the diagnostic raise errors.DSLError( node.location, f"'as_offset' is only supported for Cartesian offsets " f"(single target dimension equal to source dimension); " - f"got source '{arg_0.source.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_pretty_printer.py b/src/gt4py/next/ffront/foast_pretty_printer.py index 8b2e369501..06b4d6b228 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 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 6913b5d6be..f4c8a4fb10 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 CartesianAxisIndex, DimensionIndex + >>> class IDim(CartesianAxisIndex): ... >>> def fieldop(inp: Field[[IDim], "float64"]): ... return inp >>> @@ -232,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.value, 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.value, 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) @@ -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.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/foast_to_past.py b/src/gt4py/next/ffront/foast_to_past.py index beabe53c71..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 - >>> IDim = gtx.Dimension("I") + >>> 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 14dceb25d1..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 - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.CartesianAxisIndex): ... >>> 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 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 8cf01eb691..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 - >>> IDim = gtx.Dimension("I") + >>> 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 cd908fca14..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 - >>> IDim = gtx.Dimension("I") + >>> class IDim(gtx.CartesianAxisIndex): ... >>> @gtx.field_operator ... def copy(a: gtx.Field[[IDim], gtx.float32]) -> gtx.Field[[IDim], gtx.float32]: @@ -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,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.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( @@ -246,7 +247,8 @@ class ProgramLowering( >>> from gt4py.next import Dimension, Field >>> >>> float64 = float - >>> IDim = Dimension("IDim") + >>> 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"]): @@ -381,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.value, 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 @@ -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), 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..fd53a110f2 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 Any, Final, 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: @@ -175,13 +188,12 @@ def _scan_param_promotion( Example: -------- + >>> class I(common.CartesianAxisIndex): ... >>> _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: @@ -198,7 +210,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/field_utils.py b/src/gt4py/next/field_utils.py index 7b9fe7e68b..3026955162 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.CartesianAxisIndex): ... >>> 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.CartesianAxisIndex): ... >>> 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/fingerprinting.py b/src/gt4py/next/fingerprinting.py index f0739a102b..52a7081b1d 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 0029. + 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()), @@ -326,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, @@ -387,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/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index c8ba839a8c..7b37e35684 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.Dimension(value="_CONST_DIM", kind=common.DimensionKind.LOCAL) - - @runtime_checkable class ItIterator(Protocol): """ @@ -565,13 +559,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 == common.ConstList.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 +580,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 +600,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 @@ -729,15 +723,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.CartesianAxisIndex): ... >>> 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.CartesianAxisIndex): ... >>> j_field: LocatedField = _wrap_field( ... gtx.empty({JDim: range(3, 10)}, allocator=gtx.itir_python) ... ) @@ -753,7 +747,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 = [] @@ -895,7 +889,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 +910,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 +931,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 +968,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 +979,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 +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.value: 0}) + self._translate_named_indices({**named_indices, common.ConstList.tag: 0}) ] = value.value else: self._ndarrayfield[self._translate_named_indices(named_indices)] = value @@ -1016,7 +1010,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]: @@ -1037,13 +1033,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,10 +1150,13 @@ 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 - d, r = item[0] + 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 assert isinstance(r, core_defs.INTEGRAL_TYPES) # TODO(tehrengruber): Use a regular zero dimensional field instead. @@ -1425,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, ) @@ -1507,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.value: + if self.list_offset == common.ConstList.tag: return _ConstList( value=self.it.shift(*self.offsets, SparseTag(self.list_offset), 0).deref() ) @@ -1651,7 +1650,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 +1727,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 +1746,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))) @@ -1762,15 +1761,15 @@ 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() 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 +1823,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 +1839,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.py b/src/gt4py/next/iterator/ir.py index f024ef6168..717188c8a9 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 0029). 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/domain_utils.py b/src/gt4py/next/iterator/ir_utils/domain_utils.py index 6bd28f6dbc..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.Dimension(value=axis_literal.value, kind=axis_literal.kind)] = ( - 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: @@ -222,10 +220,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..03b7d99e71 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 @@ -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.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⟨ 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 @@ -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) def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: @@ -592,9 +592,9 @@ def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: Examples -------- - >>> IDim = common.Dimension("IDim") + >>> class IDim(common.CartesianAxisIndex): ... >>> 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))) @@ -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)) 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/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index b58f66436f..84a8008bbf 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 @@ -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 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. + 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 @@ -172,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..a7376e0d22 100644 --- a/src/gt4py/next/iterator/pretty_printer.py +++ b/src/gt4py/next/iterator/pretty_printer.py @@ -16,9 +16,10 @@ 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 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,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]: - kind = "" - if node.kind == ir.DimensionKind.HORIZONTAL: - kind = "ₕ" - elif node.kind == ir.DimensionKind.VERTICAL: - 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/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 4450d1c276..a6ae605bab 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) if isinstance(o, common.Infinity): if o is common.Infinity.POSITIVE: return itir.InfinityLiteral.POSITIVE diff --git a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py index f47d66cbe7..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,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.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 + >>> 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..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"), ... ) - >>> IDim = common.Dimension(value="IDim", kind=common.DimensionKind.HORIZONTAL) + >>> class IDim(common.CartesianAxisIndex, 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/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/prune_empty_concat_where.py b/src/gt4py/next/iterator/transforms/prune_empty_concat_where.py index 33d184e7e4..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,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.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 + >>> 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 a0c46b21de..7388cf14e2 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -24,19 +24,25 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): Example: >>> from gt4py.next import Dimension, common - >>> IDim = Dimension("IDim") - >>> JDim = Dimension("JDim") + >>> 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(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)}) >>> expr = im.call("broadcast")( ... im.ref("inp"), - ... im.make_tuple( - ... *(itir.AxisLiteral(value=dim.value, 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) >>> 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 f3696a9f80..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,8 +54,11 @@ 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.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 + >>> sys.modules[__name__].Vertex = Vertex >>> sizes = { ... "out": gtx.domain({Vertex: (0, 10), KDim: (0, 20)}), @@ -89,7 +92,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); } """ @@ -114,7 +118,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..278683429e 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.CartesianAxisIndex, 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.CartesianAxisIndex): ... >>> IHalfDim = common.flip_staggered(IDim) - >>> JDim = common.Dimension(value="JDim") + >>> 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 + >>> 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 = [] @@ -594,7 +607,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/otf/runners.py b/src/gt4py/next/otf/runners.py index f260a33061..7057e73137 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,50 @@ 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 0029), 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 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 + + 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 +254,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 +271,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/program_processors/codegens/gtfn/codegen.py b/src/gt4py/next/program_processors/codegens/gtfn/codegen.py index 3c7750184c..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 @@ -145,7 +146,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..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.value + # 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 @@ -110,10 +112,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), @@ -124,13 +132,13 @@ 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()}))" + 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 31b64c78e2..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 @@ -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=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 +118,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 +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.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 +174,35 @@ 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 +218,15 @@ 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} @@ -242,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 @@ -406,7 +423,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 = [] @@ -491,7 +508,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 ) @@ -673,12 +692,13 @@ 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), + # 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( 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.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 1eb4900233..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 @@ -58,10 +58,6 @@ ) -# 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) - - @dataclasses.dataclass(frozen=True) class ValueExpr: """ @@ -592,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 ), @@ -652,7 +648,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) ] @@ -662,7 +658,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=gtx_common.codegen_name(dim.tag))} - {offset}" for (dim, offset) in arg_expr.field_domain ) deref_node, connector_mapping = self._add_tasklet( @@ -681,7 +677,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=gtx_common.codegen_name(dim.tag)) if isinstance(index_expr, MemletExpr): self._add_input_data_edge( index_expr.dc_node, @@ -766,7 +762,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 @@ -1139,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( @@ -1152,7 +1149,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 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) index_connector = "__index" @@ -1311,10 +1312,10 @@ 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.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 @@ -1352,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) @@ -1368,7 +1369,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 @@ -1383,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( @@ -1430,7 +1432,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 @@ -1470,7 +1472,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" @@ -1482,7 +1484,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( @@ -1709,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( @@ -1914,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/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..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.value - ) + 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/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..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( @@ -77,11 +79,11 @@ 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 - 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: @@ -109,22 +111,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..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.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/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/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 09f173d3f9..90e5983527 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 0029). 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..9487baf14b 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.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. """ 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..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 - >>> I = common.Dimension(value="I") + >>> class I(common.CartesianAxisIndex): ... >>> 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.CartesianAxisIndex): ... + >>> class J(common.CartesianAxisIndex): ... >>> 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.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 - >>> JDim = common.Dimension(value="JDim") + >>> class JDim(common.CartesianAxisIndex): ... >>> 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.CartesianAxisIndex): ... + + >>> class J(common.CartesianAxisIndex): ... + + >>> class K(common.CartesianAxisIndex): ... >>> 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), @@ -903,9 +907,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..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 @@ -122,7 +126,11 @@ 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..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 @@ -219,9 +219,9 @@ 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) + dims.append(cast(common.Dimension, d)) else: raise ValueError(f"Invalid field dimensions '{dim_arg}'.") @@ -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..aea5aa7bdc 100644 --- a/tests/next_tests/artifacts/custom_named_collections.py +++ b/tests/next_tests/artifacts/custom_named_collections.py @@ -15,11 +15,21 @@ 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, + CartesianAxisIndex, + 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(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 166509e8f3..45e0c8bed6 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.CartesianAxisIndex): ... @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..65a834771f 100644 --- a/tests/next_tests/fixtures/past_common.py +++ b/tests/next_tests/fixtures/past_common.py @@ -13,9 +13,10 @@ import gt4py.next as gtx from gt4py.next import float64 - -IDim = gtx.Dimension("IDim") -JDim = gtx.Dimension("JDim") +# 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 # 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..4b73f163a9 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 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 + + __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.CartesianAxisIndex): ... + + IHalfDim = common.flip_staggered(IDim) -JDim = gtx.Dimension("JDim") + + +class JDim(gtx.CartesianAxisIndex): ... + + JHalfDim = common.flip_staggered(JDim) -KDim = gtx.Dimension("KDim", kind=gtx.DimensionKind.VERTICAL) + + +class KDim(gtx.CartesianAxisIndex, 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..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 @@ -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") @@ -108,7 +110,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 +126,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) @@ -146,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 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)}, ) e2v_np = e2v.asnumpy() @@ -167,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 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)}, ) e2v_np = e2v.asnumpy() 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..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,13 +30,22 @@ from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations -IDim = gtx.Dimension("I") +class IDim(gtx.CartesianAxisIndex): ... + + 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..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 @@ -11,12 +11,28 @@ import pytest -from gt4py.next import Dimension, DimensionKind, Field, field_operator, int32, int64, scan_operator +from gt4py.next import ( + Dimension, + CartesianAxisIndex, + 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(CartesianAxisIndex): ... + + +class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... + + @pytest.mark.parametrize( "test_case", [ @@ -45,7 +61,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 +82,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_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/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..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] -IDim = gtx.Dimension("IDim") +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 b60626acd9..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 @@ -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.CartesianAxisIndex): ... 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..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 -IDim = gtx.Dimension("IDim") +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 66a1edc189..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,8 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -I = gtx.Dimension("I") +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 5a386ee807..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,9 @@ from next_tests.unit_tests.conftest import program_processor, run_processor -I = gtx.Dimension("I") +class I(gtx.CartesianAxisIndex): ... + + 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..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,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.CartesianAxisIndex): ... + + +class JDim(gtx.CartesianAxisIndex): ... + + +class KDim(gtx.CartesianAxisIndex): ... + # 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..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 @@ -10,6 +10,7 @@ import numpy as np + try: from atlas4py import ( Config, @@ -35,14 +36,14 @@ from gt4py import next as gtx from gt4py.next.iterator import atlas_utils +# 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 -Vertex = gtx.Dimension("Vertex") -Edge = gtx.Dimension("Edge") -V2EDim = gtx.Dimension("V2E", kind=gtx.DimensionKind.LOCAL) -E2VDim = gtx.Dimension("E2V", kind=gtx.DimensionKind.LOCAL) -V2E = gtx.FieldOffset("V2E", source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset("E2V", source=Vertex, target=(Edge, E2VDim)) +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..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,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.CartesianAxisIndex): ... + + +class JDim(gtx.CartesianAxisIndex): ... + + +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_fvm_nabla.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_fvm_nabla.py index e7984f7cd3..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 @@ -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 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, + 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..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) -IDim = gtx.Dimension("IDim") +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 60aef83302..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,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.CartesianAxisIndex): ... + + +class JDim(gtx.CartesianAxisIndex): ... + 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..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 @@ -10,8 +10,11 @@ from gt4py.next import common -I = common.Dimension("I") -J = common.Dimension("J") + +class I(common.CartesianAxisIndex): ... + + +class J(common.CartesianAxisIndex): ... 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..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 -IDim = gtx.Dimension("IDim") +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 57f9ef5648..6a7cbf85de 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.CartesianAxisIndex): ... + + +class J(common.CartesianAxisIndex): ... + + +class K(common.CartesianAxisIndex): ... @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..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,6 +13,12 @@ from gt4py.next.errors import exceptions +class IDim(common.CartesianAxisIndex): ... + + +class NewDim(common.CartesianAxisIndex): ... + + 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..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,10 +18,12 @@ from gt4py.next import common, constructors from gt4py.next.common import ( Dimension, + CartesianAxisIndex, + DimensionIndex, DimensionKind, Domain, Field, - NamedIndex, + DimensionIndex, NamedRange, UnitRange, ) @@ -33,9 +35,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(CartesianAxisIndex): ... + + +class J(CartesianAxisIndex): ... + + +class I_half(CartesianAxisIndex): ... + + +class V(DimensionIndex): ... + + +class E(DimensionIndex): ... + + +class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class V2V(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class C(DimensionIndex): ... + + +class K(CartesianAxisIndex): ... + + +class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class A(DimensionIndex): ... + + +class B(DimensionIndex): ... + + +class X(CartesianAxisIndex): ... + + +class Y(CartesianAxisIndex): ... + + +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(CartesianAxisIndex): ... + + +class D1(CartesianAxisIndex): ... + + +class D2(CartesianAxisIndex): ... @pytest.fixture( @@ -67,9 +142,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 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. + """ + 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 +438,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 +461,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 +481,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 +507,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 +531,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 +555,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 +579,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 +612,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 +635,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 +659,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 +679,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 +702,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 +724,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 +750,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 +774,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 +796,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 +811,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 +827,6 @@ def test_premap_disjoint_inverse_image_raises(): def test_as_offset_1d(): # Dynamic per-point shift along I: out[i] == f[i + off[i]], full domain when all shifts in-bounds. - I = Dimension("I") Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( @@ -798,7 +843,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 +860,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 +878,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 +895,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 +911,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 +929,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 +944,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 +962,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 +971,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 +1046,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 +1055,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 +1074,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 +1115,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 +1144,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 +1346,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 +1361,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 +1382,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 +1405,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 +1417,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 +1472,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 +1493,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..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,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.CartesianAxisIndex, kind=gtx.DimensionKind.HORIZONTAL): ... + + +class VDim(gtx.CartesianAxisIndex, 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..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,9 @@ from gt4py.next.ffront.func_to_foast import FieldOperatorParser -IDim = gtx.Dimension("IDim") +class IDim(gtx.CartesianAxisIndex): ... + + 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..c00240655f 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.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 eb1927a776..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 @@ -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.CartesianAxisIndex): ... + -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.CartesianAxisIndex): ... 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) 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) 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..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,6 +54,12 @@ from gt4py.next.type_system import type_specifications as ts +class ADim(gtx.CartesianAxisIndex): ... + + +class BDim(gtx.CartesianAxisIndex): ... + + 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.CartesianAxisIndex): ... + # 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..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 @@ -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.CartesianAxisIndex): ... 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..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 -IDim = gtx.Dimension("I") +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 b6ae674312..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,8 @@ import gt4py.next.ffront.type_specifications from gt4py.next import ( Dimension, + CartesianAxisIndex, + DimensionIndex, DimensionKind, Field, FieldOffset, @@ -37,9 +39,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 0029) a redeclared `TDim` +# here would be a different dimension, where the old `Dimension("TDim")` compared equal. +TDim = cnc.TDim + + +class X(CartesianAxisIndex): ... + + +class Y(CartesianAxisIndex): ... + + +class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... + + +class ADim(CartesianAxisIndex): ... + + +class BDim(CartesianAxisIndex): ... + + +class CDim(DimensionIndex): ... + + +class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + + +class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... + + +class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class IDim(CartesianAxisIndex): ... + + +class JDim(CartesianAxisIndex): ... + + # Meaningless dimensions, used for tests. -TDim = Dimension("TDim") -SDim = Dimension("SDim") +class SDim(CartesianAxisIndex): ... def test_unpack_assign(): @@ -93,8 +136,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 +287,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 +319,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 +368,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 +380,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 +389,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 +401,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 +413,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 +425,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 +450,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 +462,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 +548,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 +558,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 +568,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 +576,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]): @@ -572,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 4aac031406..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 @@ -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.CartesianAxisIndex): ... + + 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.CartesianAxisIndex): ... + + +class K(common.CartesianAxisIndex, 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..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 ( @@ -23,10 +24,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 +51,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 @@ -62,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) @@ -76,7 +83,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 +101,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 +129,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) @@ -145,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_embedded_internals.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py index 9f61e5317e..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,6 +17,9 @@ from gt4py.next.iterator import embedded +class K(common.CartesianAxisIndex, 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..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 @@ -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.CartesianAxisIndex): ... + + 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_pretty_parser.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py index 83a16a980c..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 @@ -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, ], @@ -252,12 +252,13 @@ 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"), 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..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 @@ -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,31 @@ 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.CartesianAxisIndex): ... -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.CartesianAxisIndex, 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(): + # 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(): 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_runtime_domain.py b/tests/next_tests/unit_tests/iterator_tests/test_runtime_domain.py index bf2df06bf2..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 @@ -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.CartesianAxisIndex): ... 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..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 @@ -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, @@ -100,20 +101,16 @@ def expression_test_cases(): bool_type, ), ( - im.named_range( - itir.AxisLiteral(value="Vertex", kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 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 - ) + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1) ), ts.DomainType(dims=[Vertex]), ), @@ -131,7 +128,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 +180,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 +204,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 +213,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 +226,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 +258,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 +386,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 +409,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( @@ -442,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", 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=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( @@ -456,7 +451,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 +473,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( @@ -509,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", 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=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( @@ -540,7 +533,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 +564,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 +588,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 +615,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..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 @@ -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.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 efb050d516..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 @@ -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.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 2517c7ab55..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 @@ -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.CartesianAxisIndex, 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 b5551b6af5..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,9 +20,12 @@ ) +class I(common.CartesianAxisIndex, 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..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 @@ -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.CartesianAxisIndex): ... + + 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..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 @@ -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.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class JDim(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class KDim(common.CartesianAxisIndex, 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..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,9 @@ from gt4py.next.type_system import type_specifications as ts -IDim = common.Dimension("IDim") +class IDim(common.CartesianAxisIndex): ... + + 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..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 @@ -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.CartesianAxisIndex): ... + + +class JDim(common.CartesianAxisIndex): ... + + 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..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,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.CartesianAxisIndex): ... + + +class JDim(common.CartesianAxisIndex): ... + + +class KDim(common.CartesianAxisIndex, 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..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 @@ -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.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 b1a18ddab8..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,6 +13,9 @@ from gt4py.next.type_system import type_specifications as ts +class IDim(gtx.CartesianAxisIndex): ... + + 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..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 @@ -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.CartesianAxisIndex, 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..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 @@ -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 0029). +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..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 @@ -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.CartesianAxisIndex): ... + + +class J(gtx.CartesianAxisIndex): ... + + 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 0029). + 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..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)) -TDim = gtx.Dimension("TDim") +class TDim(gtx.CartesianAxisIndex): ... @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..0a0941d792 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) @@ -343,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 0029), 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/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..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,9 +30,11 @@ ) +class IDim(gtx.CartesianAxisIndex): ... + + @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..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 @@ -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 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 _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) @@ -376,8 +388,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_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 5a68f44961..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 @@ -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, @@ -40,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 @@ -72,48 +76,104 @@ 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}, ) 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, - __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 + }, ) @@ -308,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) @@ -493,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) @@ -755,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]: @@ -891,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) @@ -982,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) @@ -998,15 +1058,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 +1079,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 +1095,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 +1111,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 +1130,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], :][ @@ -1111,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) @@ -1153,11 +1215,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 +1234,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[ @@ -1188,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) @@ -1223,7 +1285,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 +1304,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) @@ -1268,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) @@ -1295,14 +1366,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 = [ @@ -1339,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), ) @@ -1350,14 +1421,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 = [ @@ -1396,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), ) @@ -1406,7 +1477,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 +1512,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), @@ -1459,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) @@ -1492,7 +1563,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 +1572,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) @@ -1530,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) @@ -1736,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])) @@ -1744,8 +1815,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 +1832,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")) ), @@ -1791,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), ) @@ -1818,7 +1889,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,26 +1906,26 @@ 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 ) 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, ) @@ -1937,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) @@ -1995,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) @@ -2058,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) @@ -2240,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) @@ -2335,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) @@ -2399,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/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..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 @@ -21,6 +21,13 @@ from . import util + +class boden(gtx_common.CartesianAxisIndex): ... + + +class K(gtx_common.CartesianAxisIndex, 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 fddc9dbd2e..c5daa24c9b 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -6,7 +6,10 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import itertools import operator + +import numpy as np from typing import Optional, Pattern import pytest @@ -17,6 +20,8 @@ import gt4py.next.common as common from gt4py.next.common import ( Dimension, + CartesianAxisIndex, + DimensionIndex, DimensionKind, Domain, Infinity, @@ -28,15 +33,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(CartesianAxisIndex): ... + + +class Y(CartesianAxisIndex): ... + + +class Z(CartesianAxisIndex): ... + + +class Foo(DimensionIndex): ... + + +class J(CartesianAxisIndex): ... + + +class K(CartesianAxisIndex): ... + + +class I(common.CartesianAxisIndex): ... + + +class I_half(common.CartesianAxisIndex): ... + + +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(CartesianAxisIndex): ... + + +class JDim(CartesianAxisIndex): ... + + +class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... + + IHalfDim = common.flip_staggered(IDim) @@ -393,7 +441,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) @@ -434,21 +482,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:]) @@ -461,92 +509,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, ), @@ -666,7 +714,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) @@ -675,8 +722,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) @@ -793,3 +838,157 @@ 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"), + # 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", "_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", "mod.Staggered[mod.K]"]: + 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 `'_'`. + """ + # 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, 5): + 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, 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 0029), 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") + + +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] + + @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}]") + + +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 + + +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),)) + + +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) diff --git a/tests/next_tests/unit_tests/test_constructors.py b/tests/next_tests/unit_tests/test_constructors.py index 0d6bbdeb2b..b84d5d303f 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.CartesianAxisIndex): ... + + +class J(gtx.CartesianAxisIndex): ... + + +class K(gtx.CartesianAxisIndex): ... + 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..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,6 +17,30 @@ import gt4py.storage.allocators as core_allocators +class D0(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class D1(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class D2(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... + + +class D0_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... + + +class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class D2_vertical(common.CartesianAxisIndex, kind=common.DimensionKind.VERTICAL): ... + + +class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + + +class D1_vertical(common.CartesianAxisIndex, 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.CartesianAxisIndex, 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..54e73c6a51 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.CartesianAxisIndex): ... + + @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..aad5186a85 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.CartesianAxisIndex): ... + + +class J(common.CartesianAxisIndex): ... + + @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..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,13 +12,31 @@ from gt4py.next import ( Dimension, + CartesianAxisIndex, + 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(CartesianAxisIndex): ... + + +class JDim(CartesianAxisIndex): ... + + +class KDim(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... + + +class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... + + +class TDim(CartesianAxisIndex): ... def type_info_cases() -> list[tuple[Optional[ts.TypeSpec], dict]]: @@ -61,14 +79,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 +275,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 +284,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 +306,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 +357,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 +426,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 +442,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 +457,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 +477,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..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,8 +30,11 @@ def dtype(self) -> np.dtype: return np.dtype(np.int32) -IDim = gtx.Dimension("IDim") -JDim = gtx.Dimension("JDim") +class IDim(gtx.CartesianAxisIndex): ... + + +class JDim(gtx.CartesianAxisIndex): ... + # -- 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), ), ), diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 5db040f8df..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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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]: @@ -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.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] @@ -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.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]] @@ -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.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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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 - KDim = gtx.Dimension("KDim", 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]: @@ -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.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]