diff --git a/docs/development/ADRs/next/0019-Connectivities.md b/docs/development/ADRs/next/0019-Connectivities.md index 21827941bf..3b8296459f 100644 --- a/docs/development/ADRs/next/0019-Connectivities.md +++ b/docs/development/ADRs/next/0019-Connectivities.md @@ -9,6 +9,10 @@ tags: [] - **Created**: 2024-11-08 - **Updated**: 2026-05-27 +> The `FieldOffset` part of this record is superseded by +> [ADR 0029](0029-Connectivities_As_Types.md): connectivities are declared as +> `NeighborConnectivity` classes, and offset providers are keyed by them. + The representation of Connectivities (neighbor tables, `NeighborTableOffsetProvider`) and their identifier (offset tag, `FieldOffset`, etc.) was extended and modified based on the needs of different parts of the toolchain. Here we outline the ideas for consolidating the different closely-related concepts. ## History diff --git a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md index 224f8ec289..35c2c2805b 100644 --- a/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0028-Dimensions_As_Nominal_Types.md @@ -7,7 +7,7 @@ tags: [] - **Status**: proposed - **Authors**: Enrique González Paredes (@egparedes) - **Created**: 2026-09-18 -- **Updated**: 2026-09-18 +- **Updated**: 2026-09-23 A concrete dimension becomes a **class**, and an index along it an **instance** of that class — the shape `enum.Enum` uses, where the class is the collection and @@ -77,8 +77,12 @@ disappears. 1. **Reconstruction from the IR is an import.** `common.resolve(tag)` imports the module and walks the qualname; nested declarations resolve naturally. The IR - references a Python type exactly the way `pickle` references a class. It is - memoized, because type inference calls it once per `AxisLiteral`. + references a Python type exactly the way `pickle` references a class. Where the + module path ends is memoized, because type inference calls it once per + `AxisLiteral`; the attribute walk is repeated, so a redefined declaration is + found. An + `AxisLiteral` stores only the tag: its `kind` is the resolved dimension's, so the + two cannot disagree. A purely dotted tag does not record where the module path ends and the qualname begins, so `resolve` tries the *longest importable prefix* and walks diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md new file mode 100644 index 0000000000..dec623ac50 --- /dev/null +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -0,0 +1,195 @@ +--- +tags: [] +--- + +# Connectivities as Types + +- **Status**: proposed +- **Authors**: Enrique González Paredes (@egparedes) +- **Created**: 2026-09-21 +- **Updated**: 2026-09-22 + +A neighbor connectivity is declared as a **class**, and its local dimension as a +class **nested** in it: + +```python +class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): + class Local(gtx.LocalDimensionIndex): ... + + +@gtx.field_operator +def f(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return neighbor_sum(a(V2E), axis=V2E.Local) + a(V2E[0]) +``` + +The declaration is written in DSL code, owns its local dimension, and states the +constraints a neighbor table bound to it has to satisfy. It holds no data. It builds on [ADR 0028](0028-Dimensions_As_Nominal_Types.md): the +connectivity, like a dimension, is identified by its type, and `V2E.Local` is an +ordinary dimension class with the tag `.V2E.Local`. + +## Context + +An unstructured connectivity used to be spelled by four independently authored +names that had to agree, none of them checked against the others: the +`FieldOffset` tag, the Python variable it was bound to, the local dimension's +name and the offset-provider key. The `V2EDim = Dimension("V2E")` convention made +all four equal, which hid which one each execution path actually used; the +regression tests in `test_offset_dimensions_names.py` break the convention one +name at a time. Nothing tied a local dimension to the table it indexes, so the +backends recovered that link by string equality, and the table's shape, codomain +and skip values were never checked against the `FieldOffset` declaration. + +## Decision + +### The declaration + +- `NeighborConnectivity[Origin, Codomain]` is a PEP 695 generic whose subclasses + are declarations: for each `Origin` element, a list of `Codomain` neighbors. Its + metaclass, `ConnectivityMeta`, forbids instantiation. +- The local dimension is the nested class `Local`, a subclass of + `LocalDimensionIndex`. Declaring it is required, and `NeighborConnectivity` + sets `Local.owner` to the connectivity when the class is created. A local + dimension can have at most one owner; a declaration redefined under the same + name (a re-run notebook cell) takes ownership over again, and for a local + dimension adopted rather than nested, the first declaration wins. +- A local dimension with no table, such as the coefficient axis of a fixed-size + stencil, is declared on its own: `class LsqCoeff(LocalDimensionIndex, size=3)`. + Its `owner` is `None`. A declaration can also *adopt* such a module-level local + dimension, written `Local: TypeAlias = LsqCoeff`, which then keeps its own tag. +- A connectivity can *share* another one's local dimension, + `Local: TypeAlias = C2E.Local`. + This is the flattened sparse pattern, e.g. cell-to-cell-edge (`C2CE: Cell -> CellEdge`) indexing the same neighbor axis as `C2E`, so that its results + combine with `C2E`-shaped sparse fields. The owner stays `C2E`, and the + neighbor counts and skip-value structure are the owner's. +- `max_neighbors` and `min_neighbors` are optional class keywords, not type + parameters: Python has no integer type parameters, and nothing static needs + the count. A declared count is a constraint on the bound table; an undeclared + one is taken from the table. `min_neighbors < max_neighbors` means that the + table must use skip values. +- `common.check_neighbor_table(V2E, table)` checks a table, or just its type + (which is all an ahead-of-time compilation has), against the declaration: + the domain is `(Origin, V2E.Local)`, the codomain is `Codomain`, the dtype is + integral, and the neighbor counts and skip values agree. Skip values are + checked on the table's type: a table with a `skip_value` counts as having skip + values whether or not an entry uses it. Programs run the check on the tables + they are given, see below. + +### `NeighborConnectivity` is not a `Connectivity` + +`common.Connectivity` is a *data* protocol (`ndarray`, `domain`, `asnumpy`); a +declaration holds no data. The neighbor table stays a `Connectivity` +implementation, and the declaration is only the type the table is checked +against. `Field.premap` and `Field.__call__` accept either, as they already +accepted a `FieldOffset`, which is not a `Connectivity` either. + +### `LocalDimensionIndex` subclasses `DimensionIndex` + +A separate root would force every `type[DimensionIndex]` annotation in the tree +(`ts.FieldType.dims`, `Domain`, `ConnectivityType.domain`, ...) to widen, and +would then accept local dimensions wherever a primary one is meant anyway. The +tree already distinguishes local dimensions by a runtime `kind` check, so it +keeps doing so; generic constructors whose parameter must be a primary dimension +(`NeighborConnectivity[Origin, Codomain]`, `Staggered[D]`) check it at runtime. + +### `Local` is not annotated anywhere + +Neither `NeighborConnectivity` nor `ConnectivityMeta` annotates `Local`, and +that is load-bearing: an annotation makes a declaration's `Local` a *variable* +for the checkers, so `Field[Dims[Vertex, V2E.Local], float]` is rejected by +pyright ("Variable not allowed in type expression") for a nested `Local`, and by +mypy ("not valid as a type") for an adopted or shared one. A real nested `Local` +on the base is not an option either: pyright reports an incompatible override in +every declaration. With no annotation, all three spellings are types for both +checkers, which `typing_tests/pyright_probes.py` pins for pyright and +`typing_tests/test_next.yaml` for mypy. + +The cost is that `conn.Local` is not an attribute the checkers know for a +*generic* `conn`. Library code reads it through `common.local_dimension_of(conn)` +instead, and code that has to name a local dimension generically uses a +`TypeVar` bound to `LocalDimensionIndex`. Writing an adopted or shared local as +`Local: TypeAlias = ...` (rather than a plain assignment) is what keeps mypy +treating it as a type. + +### Frontend integration + +A declaration is typed like the `FieldOffset` it replaces: `V2E.__gt_type__()` +is the `ts.OffsetType` of the derived offset `(Codomain -> (Origin, Local))`, +whose tag is the connectivity's `offset_tag`: + +- **the local dimension's tag**, `V2E.Local.tag`, for the connectivity that + declares it. This is the single string that shifts, neighbor reductions and + sparse arguments already use to find the table in the offset provider, so + existing backends need no change. +- **its own tag**, `C2CE.tag`, for a connectivity that shares another one's local + dimension, since the local dimension's tag already names the owner's table. + Shifts find the table by that tag. Reductions and sparse arguments know only + the local dimension, and take its neighbor count and skip values from a table + over it (`common.connectivity_key_over`): the owner's if bound, else the + sharer with the smallest tag. Connectivities sharing a local dimension must + therefore have the same neighbor *structure* — the same count, and a skip value + at the same positions — which is what sharing a neighbor axis means; + `check_offset_provider` enforces it for the tables it is given. + +`V2E.Local` inside DSL code types as that local dimension. The other frontend +touch points treat a declaration as the offset it replaces: grid-type deduction +(`transform_utils`, `past_to_itir`) counts it as unstructured, and embedded +`premap` accepts it. `V2E[i]` subscripts the +metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. + +### Offset providers are keyed by the declaration + +Users bind tables to declarations: + +```python +program(..., offset_provider={V2E: v2e_table, C2E: c2e_table}) +``` + +Every entry point of a program normalizes such a provider to the form the IR +uses: each declaration is replaced by its `offset_tag`. Everything below the +entry points — lowering, the backends, compiled-program caching — therefore keeps +seeing a provider keyed by strings, which is also what hand-written IR uses. + +The frontend entry points (`Program.__call__`, `FieldOperator.__call__`, +`compile`, `CompilationOptions.connectivities`) are *strict*: a string key must +be a tag, i.e. a qualified name, and a bare name such as `"V2E"` is the removed +`FieldOffset` spelling, rejected with a message pointing here. The IR-level hooks +(`embedded.context.update`, the iterator `fendef`, DaCe's `get_sdfg_conn_args`) +accept any string, because a hand-written program names its offsets itself. + +Tables are checked against their declarations (`check_offset_provider`) at every +entry point, but the result is remembered per set of bound tables, so repeated +calls cost one hash. Reading the tables — comparing the skip-value positions of +two connectivities that share a local dimension — is done only where a program is +compiled, not on the call path. A tag that +names no declared connectivity, as in hand-written IR, is not checked. + +### `FieldOffset` is removed + +`FieldOffset` and its export are gone. An unstructured connectivity is a +`NeighborConnectivity`; a Cartesian shift is `Dim + i`, which the DSL already +had; and `as_offset` takes the dimension to shift along, `as_offset(KDim, k_offsets)`, instead of a Cartesian `FieldOffset`. `scripts/python/migrate_connectivities.py` +rewrites declarations and Cartesian offset uses, and reports the provider keys +and other sites it cannot rewrite from the source alone. + +## Consequences + +- An unstructured connectivity is spelled once. The provider key, the offset tag + and the local dimension are all derived from the declaration. +- A table bound to a connectivity can be checked against its declaration. +- `V2E.Local` in DSL code is resolved from the offset type, because the type of + `V2E` is not the class. +- A declaration is fingerprinted by its name *and* its declared dimensions and + counts, so redefining it under the same name (e.g. re-running a notebook + cell) does not reuse artifacts compiled for the old declaration. +- `FieldOffset` is removed, and offset providers are keyed by declarations: a + breaking change for every unstructured program, eased by the migration script. + +## Alternatives considered + +- **The local dimension generated by the metaclass**, e.g. `V2E.Local` created + from `V2E`'s name. Type checkers cannot see a generated class, so it could not + be used in `Field[Dims[Vertex, V2E.Local], ...]`. +- **Neighbor counts as type parameters.** Python has no integer type parameters, + and a `Literal[6]` argument would add a type parameter nothing statically uses. +- **`NeighborConnectivity` as a `Connectivity` subclass.** Mixes the + declaration with the data protocol; see above. diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index 9c93fd9aa0..e663e462f3 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -23,6 +23,7 @@ Writing a new ADR is simple: - [0023 - Fingerprinting](0023-Fingerprinting.md) - [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md) - [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md) +- [0029 - Connectivities as Types](0029-Connectivities_As_Types.md) ### Frontend and Parsing #frontend diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 1e1cbfc280..8def73342a 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -218,28 +218,28 @@ edge_values = gtx.as_field([EdgeDim], np.zeros((12,))) +++ -You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _field offset_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements. +You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _connectivity_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements. To understand this transform, you can look at the edge-to-cell connectivity table `edge_to_cell_table` listed above. This table has the same shape as the output of the transform, that is, one dimension over the edges and another _local_ dimension. The table stores indices into a field over cells, the transform essentially gives you another field where the indices have been replaced with the values in the cell field at the corresponding indices. Another way to look at it is that transform uses the edge-to-cell connectivity to look up all the cell neighbors of edges, and associates the values of those neighbor cells with each edge. -You can use the field offset `E2C` below to transform a field over cells to a field over edges using the edge-to-cell connectivities: +You can use the connectivity `E2C` declared below to transform a field over cells to a field over edges using the edge-to-cell connectivities. It is declared as a class: for each edge (`EdgeDim`), a list of neighbor cells (`CellDim`), indexed by its nested local dimension `E2C.Local`: ```{code-cell} ipython3 -class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... -E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim)) +class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + class Local(gtx.LocalDimensionIndex): ... ``` -The field offset is named by its local dimension's `tag`, and the offset provider below is keyed by the same `tag`, so all three refer to one connectivity. Note that the field offset does not contain the actual connectivity table, that's provided through an _offset provider_: +Note that the declaration does not contain the actual connectivity table, that's provided through an _offset provider_, a dictionary from connectivity declarations to tables: ```{code-cell} ipython3 -E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2CDim], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) +E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2C.Local], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) ``` The field operator `nearest_cell_to_edge` below shows an example of applying this transform. There is a little twist though: the subscript in `E2C[0]` means that only the value of the first connected cell is taken, the second (if exists) is ignored. -Pay attention to the syntax where the field offset `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. +Pay attention to the syntax where the connectivity `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. ```{code-cell} ipython3 @gtx.field_operator @@ -250,7 +250,7 @@ def nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64]) -> gtx. def run_nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): nearest_cell_to_edge(cell_values, out=out) -run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2C: E2C_offset_provider}) print("0th adjacent cell's value: {}".format(edge_values.asnumpy())) ``` @@ -265,19 +265,19 @@ Running the above snippet results in the following edge field: #### Using reductions on connected mesh elements -Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2CDim) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2CDim` dimension. +Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2C.Local) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2C.Local` dimension. ```{code-cell} ipython3 @gtx.field_operator def sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64]) -> gtx.Field[Dims[EdgeDim], float64]: - # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2CDim], float64] - return neighbor_sum(cells(E2C), axis=E2CDim) + # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2C.Local], float64] + return neighbor_sum(cells(E2C), axis=E2C.Local) @gtx.program def run_sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): sum_adjacent_cells(cells, out=out) -run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2C: E2C_offset_provider}) print("sum of adjacent cells: {}".format(edge_values.asnumpy())) ``` @@ -376,13 +376,13 @@ print("where nested tuple return: {}".format(((result_1.asnumpy(), result_2.asnu #### Implementing the pseudo-laplacian -As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to define the local dimension, the field offset, and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: +As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to declare the connectivity and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: ```{code-cell} ipython3 -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=EdgeDim, target=(CellDim, C2EDim)) +class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + class Local(gtx.LocalDimensionIndex): ... -C2E_offset_provider = gtx.as_connectivity([CellDim, C2EDim], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) +C2E_offset_provider = gtx.as_connectivity([CellDim, C2E.Local], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) ``` **Weights of edge differences:** @@ -410,7 +410,7 @@ edge_weights = np.array([ [0, -1, -1], # cell 5 ], dtype=np.float64) -edge_weight_field = gtx.as_field([CellDim, C2EDim], edge_weights) +edge_weight_field = gtx.as_field([CellDim, C2E.Local], edge_weights) ``` Now you have everything to implement the pseudo-laplacian. Its field operator requires the cell field and the edge weights as inputs, and outputs a cell field of the same shape as the input. @@ -422,9 +422,9 @@ The second lines first creates a temporary field using `edge_differences(C2E)`, ```{code-cell} ipython3 @gtx.field_operator def pseudo_lap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: edges = cells(E2C[0]) # type: gtx.Field[Dims[EdgeDim], float64] - return neighbor_sum(edges(C2E) * edge_weights, axis=C2EDim) + return neighbor_sum(edges(C2E) * edge_weights, axis=C2E.Local) ``` The program itself is just a shallow wrapper over the `pseudo_lap` field operator. The significant part is how offset providers for both the edge-to-cell and cell-to-edge connectivities are supplied when the program is called: @@ -432,7 +432,7 @@ The program itself is just a shallow wrapper over the `pseudo_lap` field operato ```{code-cell} ipython3 @gtx.program def run_pseudo_laplacian(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64], + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64], out : gtx.Field[Dims[CellDim], float64]): pseudo_lap(cells, edge_weights, out=out) @@ -441,7 +441,7 @@ result_pseudo_lap = gtx.as_field([CellDim], np.zeros(shape=(6,))) run_pseudo_laplacian(cell_values, edge_weight_field, result_pseudo_lap, - offset_provider={E2CDim.tag: E2C_offset_provider, C2EDim.tag: C2E_offset_provider}) + offset_provider={E2C: E2C_offset_provider, C2E: C2E_offset_provider}) print("pseudo-laplacian: {}".format(result_pseudo_lap.asnumpy())) ``` @@ -451,7 +451,7 @@ As a closure, here is an example of chaining field operators, which is very simp ```{code-cell} ipython3 @gtx.field_operator def pseudo_laplap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: return pseudo_lap(pseudo_lap(cells, edge_weights), edge_weights) ``` diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb index b0a1980d0f..0d98021f37 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb @@ -126,7 +126,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb index 573ee6a44e..5176f48285 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb @@ -131,7 +131,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb index 2b422b1823..77d1ecf296 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb @@ -123,7 +123,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb index 85044b989f..b90958ea63 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb @@ -136,7 +136,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb index dc321f1bdd..ed332054f8 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb @@ -147,7 +147,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb index 251fe8239a..0f76b5a6b1 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb @@ -152,7 +152,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb index 30f568de6f..b557d404ea 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb @@ -293,10 +293,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E: c2e_connectivity,\n", + " V2E: v2e_connectivity,\n", + " E2V: e2v_connectivity,\n", + " E2C: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb index eaeb8c7b02..7e1a7ecc4c 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb @@ -314,10 +314,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E: c2e_connectivity,\n", + " V2E: v2e_connectivity,\n", + " E2V: e2v_connectivity,\n", + " E2C: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb b/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb index 5514c1b4f7..194d72bcc1 100644 --- a/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb +++ b/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb @@ -326,23 +326,7 @@ "execution_count": 12, "id": "ed393959", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Test successful\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/var/folders/2b/_2y31vzs4sl_7rngh2yghbpw0000gn/T/ipykernel_56692/57164705.py:4: UserWarning: Field View Program 'program_domain_where': Using Python execution, consider selecting a perfomance backend.\n", - " program_domain_where(a, b, offset_provider={\"Koff\": K})\n" - ] - } - ], + "outputs": [], "source": [ "test_domain_where()\n", "print(\"Test successful\")" diff --git a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb index b278cee26d..26973ad5b4 100644 --- a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb @@ -169,7 +169,7 @@ " kappa,\n", " dt,\n", " out=(divergence_gt4py_1, divergence_gt4py_2),\n", - " offset_provider={E2C2V.value: e2c2v_connectivity, V2E.value: v2e_connectivity},\n", + " offset_provider={E2C2V: e2c2v_connectivity, V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py_1.asnumpy(), divergence_ref_1)\n", diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index c398524538..07fb984328 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -11,7 +11,13 @@ import gt4py.next as gtx from gt4py.next.iterator.embedded import MutableLocatedField from gt4py.next import neighbor_sum, where, Dims -from gt4py.next import Dimension, DimensionIndex, DimensionKind, FieldOffset +from gt4py.next import ( + Dimension, + DimensionIndex, + DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, +) from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( run_gtfn as gtfn_cpu, @@ -389,31 +395,36 @@ class E(DimensionIndex): ... class K(DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E(NeighborConnectivity[C, E]): + class Local(LocalDimensionIndex): ... -C2E = FieldOffset(C2EDim.tag, source=E, target=(C, C2EDim)) +C2EDim = C2E.Local -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(NeighborConnectivity[V, E]): + class Local(LocalDimensionIndex): ... -V2E = FieldOffset(V2EDim.tag, source=E, target=(V, V2EDim)) +V2EDim = V2E.Local -class E2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(NeighborConnectivity[E, V]): + class Local(LocalDimensionIndex): ... -E2V = FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) +E2VDim = E2V.Local -class E2CDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C(NeighborConnectivity[E, C]): + class Local(LocalDimensionIndex): ... -E2C = FieldOffset(E2CDim.tag, source=C, target=(E, E2CDim)) +E2CDim = E2C.Local -class E2C2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2V(NeighborConnectivity[E, V]): + class Local(LocalDimensionIndex): ... -E2C2V = FieldOffset(E2C2VDim.tag, source=V, target=(E, E2C2VDim)) +E2C2VDim = E2C2V.Local diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index db8f370abc..2cde0190ff 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -273,10 +273,11 @@ "metadata": {}, "outputs": [], "source": [ - "class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...\n", + "class E2C(gtx.NeighborConnectivity[Edge, Cell]):\n", + " class Local(gtx.LocalDimensionIndex): ...\n", "\n", "\n", - "E2C = gtx.FieldOffset(E2CDim.tag, source=Cell, target=(Edge, E2CDim))" + "E2CDim = E2C.Local" ] }, { @@ -320,7 +321,7 @@ " nearest_cell_to_edge(cell_field, out=edge_field)\n", "\n", "\n", - "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", + "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={E2C: E2C_offset_provider})\n", "\n", "print(\"0th adjacent cell's value: {}\".format(edge_field.asnumpy()))" ] @@ -395,7 +396,7 @@ " sum_adjacent_cells(cell_field, out=edge_field)\n", "\n", "\n", - "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", + "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={E2C: E2C_offset_provider})\n", "\n", "print(\"sum of adjacent cells: {}\".format(edge_field.asnumpy()))" ] diff --git a/noxfile.py b/noxfile.py index c8091de09e..31327f076d 100755 --- a/noxfile.py +++ b/noxfile.py @@ -349,6 +349,9 @@ def test_typing_exports(session: nox.Session) -> None: "typing_tests", *session.posargs, ) + # A second checker, on code that must type-check for a downstream user: mypy and pyright + # disagree about what counts as a type, which the mypy-only cases above cannot catch. + session.run("pyright", "--project", "typing_tests", "typing_tests/pyright_probes.py") # -- DaCe codegen determinism check -- diff --git a/pyproject.toml b/pyproject.toml index 64f52fdb8b..3a5ea4af72 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -64,6 +64,7 @@ typing = [ typing_exports = [ # to test typing with gt4py in downstream code {include-group = "typing"}, + 'pyright>=1.1.400', # the second checker: it disagrees with mypy about what counts as a type 'types-six>=1.17.0.20251009', # can not let mypy auto-install types as that leads to unexpected stderr output (which means test failure) 'pytest-mypy-plugins>=4.0.0', # pytest plugin for running mypy on code snippets "xarray>=2024.1.0" # one of the regression tests requires xarray @@ -262,8 +263,6 @@ markers = [ 'uses_ir_if_stmts', 'uses_lift: tests that require backend support for lift builtin function', 'uses_negative_modulo: tests that require backend support for modulo on negative numbers', - 'uses_offset_tag_differing_from_local_dim: tests that shift with an offset tag that differs from the name of its local dimension', - 'uses_offset_tag_differing_from_local_dim_in_reduction: tests that reduce over an offset whose tag differs from the name of its local dimension', 'uses_origin: tests that require backend support for domain origin', 'uses_reduce_with_lambda: tests that use lambdas as reduce functions', 'uses_reduction_with_only_sparse_fields: tests that require backend support for with sparse fields', diff --git a/scripts/python/migrate_connectivities.py b/scripts/python/migrate_connectivities.py new file mode 100644 index 0000000000..d3523b1887 --- /dev/null +++ b/scripts/python/migrate_connectivities.py @@ -0,0 +1,473 @@ +#!/usr/bin/env -S uv run -q --frozen --isolated --python 3.12 --group scripts python3 +# +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +""" +Migrate gt4py.next user code to dimension and connectivity classes (ADRs 0028, 0029). + +Rewrites module-level declarations and the uses of Cartesian offsets: + + KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) + E2CDim = gtx.Dimension("E2C", gtx.DimensionKind.LOCAL) + E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim, E2CDim)) + Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) + ... a(Koff[1]) ... as_offset(Koff, k_field) ... + +becomes + + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class E2CDim(gtx.LocalDimensionIndex): ... + class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + Local: typing.TypeAlias = E2CDim + ... a(KDim + 1) ... as_offset(KDim, k_field) ... + +A connectivity adopts its existing local dimension (`Local: TypeAlias = E2CDim`), so the names already used +for local dimensions keep working, and so do offsets that share a local dimension (`C2CE` +with `C2EDim`). What cannot be rewritten from the source alone is reported instead: offset +providers keyed by strings, which become keyed by the connectivity (`{E2C: table}`), `.value` +on a dimension (now `.tag`), and `isinstance` checks against `Dimension`. + +The output is not formatted; run `ruff format` on the changed files afterwards. +""" + +from __future__ import annotations + +import ast +import dataclasses +import difflib +import pathlib +import re +from collections.abc import Iterable, Iterator + +import typer + + +cli = typer.Typer(no_args_is_help=True, name="migrate-connectivities", help=__doc__) + + +@dataclasses.dataclass(frozen=True) +class Edit: + """Replace source lines `[start, end)` (0-based) with `text`.""" + + start: int + end: int + text: str + + +@dataclasses.dataclass +class Module: + path: pathlib.Path + source: str + tree: ast.Module + edits: list[Edit] = dataclasses.field(default_factory=list) + notes: list[str] = dataclasses.field(default_factory=list) + #: Names of `gt4py.next` the migrated declarations use unqualified, to be imported. + needed: set[str] = dataclasses.field(default_factory=set) + #: Whether the migrated declarations need `typing` imported (for `Local: TypeAlias = ...`). + needed_typing: bool = False + + @property + def lines(self) -> list[str]: + return self.source.splitlines(keepends=True) + + def segment(self, node: ast.AST) -> str: + segment = ast.get_source_segment(self.source, node) + assert segment is not None + return segment + + def note(self, node: ast.stmt | ast.expr, message: str) -> None: + self.notes.append(f"{self.path}:{node.lineno}: {message}") + + +def _callee_name(call: ast.Call) -> tuple[str, str] | None: + """`(prefix, name)` of a call to `[prefix.]name`, e.g. `("gtx.", "Dimension")`.""" + match call.func: + case ast.Name(id=name): + return "", name + case ast.Attribute(value=value, attr=name): + return f"{ast.unparse(value)}.", name + return None + + +def _keyword(call: ast.Call, name: str, position: int) -> ast.expr | None: + for keyword in call.keywords: + if keyword.arg == name: + return keyword.value + return call.args[position] if len(call.args) > position else None + + +_DECLARATION_CALLEES = ("Dimension", "FieldOffset") + + +def _imported_aliases(module: Module) -> dict[str, str]: + """Local names of imported `Dimension` / `FieldOffset`, e.g. `{"FO": "FieldOffset"}`.""" + return { + alias.asname or alias.name: alias.name + for statement in ast.walk(module.tree) + if isinstance(statement, ast.ImportFrom) + for alias in statement.names + if alias.name in _DECLARATION_CALLEES + } + + +def _declarations(module: Module) -> Iterator[tuple[ast.Assign, str, ast.Call, str, str]]: + """Module-level `name = [prefix.]Dimension(...)` / `FieldOffset(...)` statements.""" + aliases = _imported_aliases(module) + for statement in module.tree.body: + if ( + isinstance(statement, ast.Assign) + and len(statement.targets) == 1 + and isinstance(statement.targets[0], ast.Name) + and isinstance(statement.value, ast.Call) + and (callee := _callee_name(statement.value)) is not None + ): + prefix, name = callee + name = aliases.get(name, name) if prefix == "" else name + if name in _DECLARATION_CALLEES: + yield statement, statement.targets[0].id, statement.value, prefix, name + + +def _replace(module: Module, statement: ast.stmt, text: str) -> None: + assert statement.end_lineno is not None + module.edits.append(Edit(statement.lineno - 1, statement.end_lineno, text)) + + +def _migrate_declarations(module: Module, cartesian: dict[str, str]) -> None: + for statement, name, call, prefix, kind_of_call in _declarations(module): + if kind_of_call == "Dimension": + kind = _keyword(call, "kind", 1) + kind_src = module.segment(kind) if kind is not None else None + if kind_src is not None and kind_src.split(".")[-1] == "LOCAL": + base = "LocalDimensionIndex" + text = f"class {name}({prefix}{base}): ...\n" + elif kind_src is None: + base = "DimensionIndex" + text = f"class {name}({prefix}{base}): ...\n" + else: + base = "DimensionIndex" + text = f"class {name}({prefix}{base}, kind={kind_src}): ...\n" + if not prefix: + module.needed.add(base) + _replace(module, statement, text) + continue + + source, target = _keyword(call, "source", 1), _keyword(call, "target", 2) + if source is None or not isinstance(target, ast.Tuple): + module.note(statement, f"'{name}': unrecognized 'FieldOffset' arguments, not migrated.") + continue + if len(target.elts) == 2: + origin, local = (module.segment(element) for element in target.elts) + text = ( + f"class {name}({prefix}NeighborConnectivity[{origin}, {module.segment(source)}]):\n" + # NOTE: `TypeAlias`, not a plain assignment: it is what keeps the adopted local + # dimension a *type* for mypy (see ADR 0029). + f" Local: typing.TypeAlias = {local}\n" + ) + module.needed_typing = True + if not prefix: + module.needed.add("NeighborConnectivity") + _replace(module, statement, text) + elif len(target.elts) == 1 and module.segment(target.elts[0]) == module.segment(source): + # A Cartesian offset has no declaration any more: `Off[i]` is `Dim + i`. + cartesian[name] = module.segment(source) + _replace(module, statement, "") + else: + module.note(statement, f"'{name}': a cross-dimension offset has no class equivalent.") + + +def _migrate_imports(module: Module) -> None: + """Drop imports of the removed `FieldOffset`; import the class names used unqualified.""" + fieldoffset_names = { + local for local, name in _imported_aliases(module).items() if name == "FieldOffset" + } + imported = { + alias.asname or alias.name + for statement in ast.walk(module.tree) + if isinstance(statement, ast.ImportFrom) + for alias in statement.names + } + missing = sorted(module.needed - imported) + typing_missing = module.needed_typing and not any( + isinstance(statement, ast.Import) and any(a.name == "typing" for a in statement.names) + for statement in ast.walk(module.tree) + ) + added = False + for statement in module.tree.body: + if not isinstance(statement, ast.ImportFrom): + continue + names = {alias.asname or alias.name for alias in statement.names} + drops = names & fieldoffset_names + adds_here = bool(missing) and not added and bool(names & {"Dimension", "DimensionKind"}) + if not (drops or adds_here): + continue + kept = [ + ast.unparse(alias) + for alias in statement.names + if (alias.asname or alias.name) not in fieldoffset_names + ] + text = ( + f"from {'.' * statement.level}{statement.module or ''} import {', '.join(kept)}\n" + if kept + else "" + ) + if adds_here: + text += f"from gt4py.next import {', '.join(missing)}\n" + added = True + _replace(module, statement, text) + if typing_missing: + # `Local: TypeAlias = ...` needs it; the first import statement is a safe place + for statement in module.tree.body: + if isinstance(statement, (ast.Import, ast.ImportFrom)): + module.edits.append( + Edit(statement.lineno - 1, statement.lineno - 1, "import typing\n") + ) + break + else: + module.notes.append( + f"{module.path}: import 'typing', used by the migrated declarations." + ) + if missing and not added: + module.notes.append( + f"{module.path}: import {', '.join(missing)} from 'gt4py.next', used by the migrated" + " declarations." + ) + + +class _CartesianUses(ast.NodeVisitor): + """Rewrite the uses of removed Cartesian offsets, and note the ones it cannot.""" + + def __init__(self, module: Module, cartesian: dict[str, str]) -> None: + self.module = module + self.cartesian = cartesian + #: (line, start column, end column, replacement) of single-line expression rewrites + self.rewrites: list[tuple[int, int, int, str]] = [] + self.handled: set[int] = set() + #: names bound in the enclosing function scopes, which shadow a removed offset + self.shadowed: list[set[str]] = [] + + def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> None: + arguments = node.args + bound = { + argument.arg + for argument in (*arguments.posonlyargs, *arguments.args, *arguments.kwonlyargs) + } | { + name.id + for name in ast.walk(node) + if isinstance(name, ast.Name) and isinstance(name.ctx, ast.Store) + } + self.shadowed.append(bound) + self.generic_visit(node) + self.shadowed.pop() + + visit_AsyncFunctionDef = visit_FunctionDef + visit_Lambda = visit_FunctionDef + + def _rewrite(self, node: ast.expr, text: str) -> None: + assert node.end_lineno is not None and node.end_col_offset is not None + if node.lineno != node.end_lineno: + self.module.note(node, f"multi-line expression; rewrite by hand as '{text}'.") + return + self.rewrites.append((node.lineno - 1, node.col_offset, node.end_col_offset, text)) + + def _dimension_of(self, node: ast.expr) -> str | None: + """The dimension replacing `Off` or `module.Off`, qualified like the offset was.""" + match node: + case ast.Name(id=name) if name in self.cartesian and not any( + name in bound for bound in self.shadowed + ): + return self.cartesian[name] + case ast.Attribute(value=value, attr=name) if name in self.cartesian: + return f"{self.module.segment(value)}.{self.cartesian[name]}" + return None + + def visit_Subscript(self, node: ast.Subscript) -> None: + if (dim := self._dimension_of(node.value)) is not None: + match node.slice: + case ast.Constant(value=int() as index): + text = f"{dim} + {index}" if index >= 0 else f"{dim} - {-index}" + case ast.UnaryOp(op=ast.USub(), operand=ast.Constant(value=int() as index)): + text = f"{dim} - {index}" + case _: + text = f"{dim} + ({self.module.segment(node.slice)})" + self._rewrite(node, text) + self.handled.add(id(node.value)) + self.generic_visit(node) + + def visit_Name(self, node: ast.Name) -> None: + if id(node) not in self.handled and (dim := self._dimension_of(node)) is not None: + # e.g. `as_offset(Koff, field)`, which now takes the dimension + self._rewrite(node, dim) + + def visit_Attribute(self, node: ast.Attribute) -> None: + if id(node) not in self.handled and (dim := self._dimension_of(node)) is not None: + self._rewrite(node, dim) + else: + self.generic_visit(node) + + +def _migrate_cartesian_uses(module: Module, cartesian: dict[str, str]) -> None: + if not cartesian: + return + visitor = _CartesianUses(module, cartesian) + import_lines: set[int] = set() + for statement in ast.walk(module.tree): + if isinstance(statement, ast.ImportFrom) and any( + alias.name in cartesian for alias in statement.names + ): + import_lines.update(range(statement.lineno - 1, statement.end_lineno or 0)) + names = [ + cartesian.get(alias.name, alias.name) if alias.asname is None else alias.name + for alias in statement.names + ] + unique = list(dict.fromkeys(names)) + _replace( + module, + statement, + " " * statement.col_offset + + f"from {'.' * statement.level}{statement.module or ''} import {', '.join(unique)}\n", + ) + for statement in module.tree.body: + if not isinstance(statement, (ast.Import, ast.ImportFrom)): + visitor.visit(statement) + match statement: + case ast.Assign(targets=[ast.Name(id="__all__")], value=ast.List(elts=all_names)): + for entry in all_names: + if isinstance(entry, ast.Constant) and entry.value in cartesian: + module.note(entry, f"'__all__' lists the removed offset '{entry.value}'.") + + lines = module.lines + for line, start, end, text in sorted(visitor.rewrites, reverse=True): + if line in import_lines: + continue + lines[line] = lines[line][:start] + text + lines[line][end:] + module.source = "".join(lines) + + +_PROVIDER_KEY_RE = re.compile(r"""(?P["'])(?P[A-Za-z_]\w*)(?P=quote)\s*:""") + + +def _report( + module: Module, + offset_keys: dict[str, str], + cartesian: dict[str, str], + dimension_names: set[str], +) -> None: + for number, line in enumerate(module.source.splitlines(), start=1): + for match in _PROVIDER_KEY_RE.finditer(line): + if (connectivity := offset_keys.get(match["name"])) is None: + continue + if connectivity in cartesian: + message = ( + f"offset-provider key '{match['name']}': remove the entry, a Cartesian shift" + f" ('{cartesian[connectivity]} + i') needs none." + ) + else: + message = ( + f"offset-provider key '{match['name']}' is keyed by the connectivity class" + f" now, e.g. '{{{connectivity}: table}}'." + ) + module.notes.append(f"{module.path}:{number}: {message}") + for node in ast.walk(module.tree): + match node: + case ast.Attribute( + value=ast.Name(id=name) | ast.Attribute(attr=name), attr="value" + ) if name in dimension_names: + module.note(node, f"'{name}.value': a dimension's name is '{name}.tag' now.") + case ast.Call( + func=ast.Name(id="isinstance"), args=[_, ast.Attribute(attr="Dimension")] + ): + module.note( + node, + "'isinstance(..., Dimension)': a dimension is a class now; use" + " 'isinstance(obj, gt4py.next.common.DimensionMeta)'.", + ) + + +def _apply(module: Module) -> str: + lines = module.lines + for edit in sorted(module.edits, key=lambda edit: edit.start, reverse=True): + lines[edit.start : edit.end] = [edit.text] if edit.text else [] + return "".join(lines) + + +def migrate(sources: dict[pathlib.Path, str]) -> tuple[dict[pathlib.Path, str], list[str]]: + """Migrate the given modules; return their new sources and what is left to do by hand.""" + modules = [ + Module(path=path, source=source, tree=ast.parse(source)) for path, source in sources.items() + ] + cartesian: dict[str, str] = {} + #: offset-provider keys that name a `FieldOffset`: its variable name and its tag, if different + offset_keys: dict[str, str] = {} + dimension_names: set[str] = set() + for module in modules: + for _, name, call, _, kind_of_call in _declarations(module): + if kind_of_call == "Dimension": + dimension_names.add(name) + continue + offset_keys[name] = name + if call.args and isinstance(tag := call.args[0], ast.Constant): + offset_keys[str(tag.value)] = name + _migrate_declarations(module, cartesian) + _migrate_imports(module) + + results: dict[pathlib.Path, str] = {} + notes: list[str] = [] + for module in modules: + _report(module, offset_keys, cartesian, dimension_names) + migrated = _apply(module) + # Cartesian uses are rewritten on the migrated text, re-parsed, so that line numbers + # refer to what the declaration edits left. + second = Module(path=module.path, source=migrated, tree=ast.parse(migrated)) + _migrate_cartesian_uses(second, cartesian) + results[module.path] = _apply(second) + notes += module.notes + second.notes + return results, notes + + +def _python_files(paths: Iterable[pathlib.Path]) -> Iterator[pathlib.Path]: + for path in paths: + if path.is_dir(): + yield from sorted(path.rglob("*.py")) + else: + yield path + + +@cli.command() +def run( + paths: list[pathlib.Path], + write: bool = typer.Option(False, "--write", help="Rewrite files instead of printing a diff."), +) -> None: + """Migrate the Python files in PATHS (directories are searched recursively).""" + sources = {path: path.read_text() for path in _python_files(paths)} + results, notes = migrate(sources) + for path, new in results.items(): + if new == sources[path]: + continue + if write: + path.write_text(new) + typer.echo(f"migrated {path}") + else: + typer.echo( + "".join( + difflib.unified_diff( + sources[path].splitlines(keepends=True), + new.splitlines(keepends=True), + fromfile=str(path), + tofile=str(path), + ) + ) + ) + if notes: + typer.echo("\nLeft to migrate by hand (line numbers of the original files):", err=True) + for note in notes: + typer.echo(f" {note}", err=True) + + +if __name__ == "__main__": + cli() diff --git a/scripts/tests/python/test_migrate_connectivities.py b/scripts/tests/python/test_migrate_connectivities.py new file mode 100644 index 0000000000..4d196eeeba --- /dev/null +++ b/scripts/tests/python/test_migrate_connectivities.py @@ -0,0 +1,174 @@ +# +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause +# + +"""Tests for the ``migrate_connectivities`` dev script.""" + +from __future__ import annotations + +import pathlib +import textwrap + +import migrate_connectivities + + +DIMENSIONS = textwrap.dedent( + """\ + import gt4py.next as gtx + + KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) + EdgeDim = gtx.Dimension("Edge") + CellDim = gtx.Dimension("Cell") + CEDim = gtx.Dimension("CE") + E2CDim = gtx.Dimension("E2C", gtx.DimensionKind.LOCAL) + C2EDim = gtx.Dimension("C2E", gtx.DimensionKind.LOCAL) + E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim, E2CDim)) + C2E = gtx.FieldOffset("C2E", source=EdgeDim, target=(CellDim, C2EDim)) + C2CE = gtx.FieldOffset("C2CE", source=CEDim, target=(CellDim, C2EDim)) + Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) + """ +) + +STENCIL = textwrap.dedent( + """\ + import gt4py.next as gtx + from gt4py.next.ffront.experimental import as_offset + + from pkg import dimension as dims + from pkg.dimension import E2C, KDim, Koff + + + def stencil(a, k): + b = a(Koff[1]) + a(Koff[-1]) + a(Koff[k]) + a(dims.Koff[1]) + return b(as_offset(Koff, k)) + b(as_offset(dims.Koff, k)) + + + def run(prog, grid): + prog(offset_provider={"E2C": grid.e2c, "Koff": KDim}) + return KDim.value + """ +) + + +def _migrate(**sources: str) -> tuple[dict[str, str], list[str]]: + results, notes = migrate_connectivities.migrate( + {pathlib.Path(name): source for name, source in sources.items()} + ) + return {str(path): source for path, source in results.items()}, notes + + +def test_declarations(): + results, _ = _migrate(dimension=DIMENSIONS) + + assert results["dimension"] == textwrap.dedent( + """\ + import typing + import gt4py.next as gtx + + class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + class EdgeDim(gtx.DimensionIndex): ... + class CellDim(gtx.DimensionIndex): ... + class CEDim(gtx.DimensionIndex): ... + class E2CDim(gtx.LocalDimensionIndex): ... + class C2EDim(gtx.LocalDimensionIndex): ... + class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + Local: typing.TypeAlias = E2CDim + class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + Local: typing.TypeAlias = C2EDim + class C2CE(gtx.NeighborConnectivity[CellDim, CEDim]): + Local: typing.TypeAlias = C2EDim + """ + ) + + +def test_declarations_run(): + results, _ = _migrate(dimension=DIMENSIONS) + namespace: dict = {"__name__": "migrated_dimension"} + exec(results["dimension"], namespace) + + assert namespace["C2E"].Local is namespace["C2EDim"] + assert namespace["C2EDim"].owner is namespace["C2E"] + # `C2CE` shares `C2E`'s local dimension and is named by its own tag + assert namespace["C2CE"].Local is namespace["C2EDim"] + assert namespace["C2CE"].offset_tag == namespace["C2CE"].tag + + +def test_cartesian_offset_uses_across_modules(): + results, _ = _migrate(dimension=DIMENSIONS, stencil=STENCIL) + + stencil = results["stencil"] + assert "from pkg.dimension import E2C, KDim\n" in stencil + assert "a(KDim + 1) + a(KDim - 1) + a(KDim + (k)) + a(dims.KDim + 1)" in stencil + assert "b(as_offset(KDim, k)) + b(as_offset(dims.KDim, k))" in stencil + assert "Koff" not in stencil.replace('"Koff"', "") + + +def test_what_is_left_is_reported(): + _, notes = _migrate(dimension=DIMENSIONS, stencil=STENCIL) + + assert any("offset-provider key 'E2C'" in note for note in notes) + assert any("offset-provider key 'Koff'" in note for note in notes) + assert any("'KDim.value'" in note for note in notes) + + +BARE = textwrap.dedent( + """\ + from gt4py.next import Dimension, DimensionKind, FieldOffset as FO + + LOCAL = DimensionKind.LOCAL + Vertex = Dimension("Vertex") + Edge = Dimension("Edge") + V2EDim = Dimension("V2E", LOCAL) + V2E = FO("V2E_TAG", source=Edge, target=(Vertex, V2EDim)) + """ +) + + +def test_unqualified_names_and_aliases(): + results, notes = _migrate(bare=BARE, user='table = {"V2E_TAG": t}\n') + migrated = results["bare"] + + assert "FO" not in migrated + assert "from gt4py.next import Dimension, DimensionKind\n" in migrated + assert ( + "from gt4py.next import DimensionIndex, LocalDimensionIndex, NeighborConnectivity\n" + in migrated + ) + assert "class V2EDim(LocalDimensionIndex): ..." in migrated + assert "class V2E(NeighborConnectivity[Vertex, Edge]):" in migrated + assert " Local: typing.TypeAlias = V2EDim" in migrated + namespace: dict = {"__name__": "migrated_bare"} + exec(migrated, namespace) + assert namespace["V2E"].Local is namespace["V2EDim"] + # a key spelled with the offset's tag is reported, naming the class + assert any("'V2E_TAG'" in note and "{V2E: table}" in note for note in notes) + + +def test_shadowed_names_and_all_are_left_alone(): + source = textwrap.dedent( + """\ + from pkg.dimension import KDim, Koff + + __all__ = ["Koff"] + + + def helper(Koff): + return Koff + 1 + + + def uses(a): + return a(Koff[1]) + a(dims.KDim.value) + """ + ) + results, notes = _migrate(dimension=DIMENSIONS, user=source) + + assert "return Koff + 1" in results["user"] + assert "a(KDim + 1)" in results["user"] + assert any("'__all__' lists the removed offset 'Koff'" in note for note in notes) + assert any("'KDim.value'" in note for note in notes) diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index e665024d7d..1fd3b4d806 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -31,12 +31,15 @@ Domain, Field, GridType, + LocalDimensionIndex, + NeighborConnectivity, Staggered, UnitRange, as_non_staggered, domain, flip_staggered, is_staggered, + local_dimension_of, resolve, unit_range, ) @@ -47,7 +50,6 @@ from .ffront import fbuiltins from .ffront.decorator import field_operator, program, scan_operator from .ffront.fbuiltins import ( - FieldOffset, IndexType, abs, # noqa: A004 # shadowing arccos, @@ -120,6 +122,8 @@ "Dimension", "DimensionIndex", "DimensionKind", + "LocalDimensionIndex", + "NeighborConnectivity", "Staggered", "resolve", "Dims", @@ -132,6 +136,7 @@ "unit_range", "UnitRange", "is_staggered", + "local_dimension_of", "flip_staggered", "as_non_staggered", # from constructors @@ -143,7 +148,6 @@ "as_field", "as_connectivity", # from ffront - "FieldOffset", "field_operator", "program", "scan_operator", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 89014c25c0..64c6e58a3c 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -16,6 +16,7 @@ import functools import importlib import math +import numbers import re import sys import types @@ -283,6 +284,14 @@ def __init_subclass__(cls, /, kind: Optional[DimensionKind] = None, **kwargs: An ) if kind is not None: cls.kind = kind + if cls.kind is DimensionKind.LOCAL and not any( + "_local_dimension_root" in base.__dict__ for base in cls.__mro__ + ): + raise TypeError( + f"'{cls.__qualname__}': a local dimension is declared by subclassing" + " 'LocalDimensionIndex', or as the nested 'Local' class of a" + " 'NeighborConnectivity', not with 'kind=DimensionKind.LOCAL'." + ) def __init__(self, value: int) -> None: self.value = value @@ -343,7 +352,6 @@ def staggered_base_tag(tag: Tag) -> Optional[Tag]: return match["base"] if (match := _STAGGERED_TAG_RE.match(tag)) is not None else None -@functools.cache def resolve(tag: Tag) -> Dimension: """ Return the dimension class a tag names, by importing it. @@ -356,6 +364,10 @@ def resolve(tag: Tag) -> Dimension: so the longest importable prefix wins and the remainder is walked as attributes. A collision would need a module path and an attribute chain to have the same spelling. + Only where the module path ends is memoized; the attribute walk is repeated on every call, so + a declaration redefined under the same name (e.g. by re-running a notebook cell) resolves to + the new class. + Parametrized dimensions such as `Staggered[K]` have no importable qualname; their tag has the form `[]` and is resolved by subscripting the owner, which goes through its intern table and so returns the identical class. @@ -381,25 +393,86 @@ def resolve(tag: Tag) -> Dimension: ) return owner[resolve(match["base"])] # type: ignore[index] # a StaggeredMeta, checked + obj = _import_qualified_name(tag) + if not isinstance(obj, DimensionMeta): + raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") + return cast(Dimension, obj) + + +def _import_qualified_name_or_none(tag: Tag) -> Any: + """`_import_qualified_name`, returning `None` for a name that does not resolve.""" + if (split := _split_qualified_name_or_none(tag)) is None: + return None + module_name, attrs = split + obj: Any = sys.modules.get(module_name) or importlib.import_module(module_name) + for attr in attrs: + if (obj := getattr(obj, attr, None)) is None: + return None + return obj + + +def resolve_loaded(tag: Tag) -> Optional[Dimension]: + """ + Return the dimension a tag names if its module is already loaded, else `None`. + + Like `resolve`, but never imports: for code that must not have import side effects, such as + printing IR. + """ + if (match := _STAGGERED_TAG_RE.match(tag)) is not None: + owner, base = resolve_loaded(match["owner"]), resolve_loaded(match["base"]) + return owner[base] if owner is not None and base is not None else None # type: ignore[index] # parametrized dimension + parts = tag.split(".") + for split in range(len(parts) - 1, 0, -1): + if (obj := sys.modules.get(".".join(parts[:split]))) is None: + continue + for attr in parts[split:]: + obj = getattr(obj, attr, None) + return obj if isinstance(obj, DimensionMeta) else None + return None + + +def _import_qualified_name(tag: Tag) -> Any: + """Import the object a dotted qualified name refers to; see `resolve`.""" + module_name, attrs = _split_qualified_name(tag) + obj: Any = sys.modules.get(module_name) or importlib.import_module(module_name) + for attr in attrs: + try: + obj = getattr(obj, attr) + except AttributeError as ex: + raise ValueError( + f"Cannot resolve tag '{tag}': '{module_name}' has no attribute '{'.'.join(attrs)}'." + ) from ex + return obj + + +@functools.cache +def _split_qualified_name_or_none(tag: Tag) -> Optional[tuple[str, tuple[str, ...]]]: + """ + `_split_qualified_name`, returning `None` instead of raising. + + Separate and memoized so that a string that is not a qualified name -- an offset-provider key + of hand-written IR, say -- costs one import attempt in total, not one per call. A module whose + import *fails* other than by not being found is reported, not cached away. + """ + try: + return _split_qualified_name(tag) + except ValueError: + return None + + +@functools.cache +def _split_qualified_name(tag: Tag) -> tuple[str, tuple[str, ...]]: + """Split a dotted name at its longest importable module prefix: `(module, attributes)`.""" parts = tag.split(".") for split in range(len(parts), 0, -1): + module_name = ".".join(parts[:split]) try: - obj: Any = importlib.import_module(".".join(parts[:split])) + importlib.import_module(module_name) except ImportError: continue - for attr in parts[split:]: - try: - obj = getattr(obj, attr) - except AttributeError as ex: - raise ValueError( - f"Cannot resolve dimension tag '{tag}': '{'.'.join(parts[:split])}' has" - f" no attribute '{attr}'." - ) from ex - if not isinstance(obj, DimensionMeta): - raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") - return cast(Dimension, obj) + return module_name, tuple(parts[split:]) raise ValueError( - f"Cannot resolve dimension tag '{tag}': no importable module prefix. A dimension" + f"Cannot resolve tag '{tag}': no importable module prefix. A dimension or connectivity" " referenced from the IR must be declared at module level in an importable module." ) @@ -1027,7 +1100,7 @@ def asnumpy(self) -> np.ndarray: ... def as_scalar(self) -> core_defs.ScalarT: ... @abc.abstractmethod - def premap(self, index_field: Connectivity | fbuiltins.FieldOffset) -> Field: ... + def premap(self, index_field: Connectivity | type[NeighborConnectivity]) -> Field: ... @abc.abstractmethod def restrict(self, item: AnyIndexSpec) -> Self: ... @@ -1035,8 +1108,8 @@ def restrict(self, item: AnyIndexSpec) -> Self: ... @abc.abstractmethod def __call__( self, - index_field: Connectivity | fbuiltins.FieldOffset, - *args: Connectivity | fbuiltins.FieldOffset, + index_field: Connectivity | type[NeighborConnectivity], + *args: Connectivity | type[NeighborConnectivity], ) -> Field: ... @abc.abstractmethod @@ -1219,7 +1292,10 @@ def has_skip_values(self) -> bool: @dataclasses.dataclass(frozen=True) class NeighborConnectivityType(ConnectivityType): - # TODO(havogt): refactor towards encoding this information in the local dimensions of the ConnectivityType.domain + # NOTE: partly encoded in the local dimension since ADR 0029: a `LocalDimensionIndex` carries + # `max_neighbors` / `min_neighbors` where the declaration states them, and this record is + # checked against them (`check_neighbor_table`). It stays the *bound* count, which a + # declaration may leave to the table. max_neighbors: int @property @@ -1418,8 +1494,16 @@ def is_neighbor_table(obj: Any) -> TypeGuard[NeighborTable]: OffsetProviderTypeElem: TypeAlias = NeighborConnectivityType # Note: `OffsetProvider` and `OffsetProviderType` should not be accessed directly, # use the `get_offset` and `get_offset_type` functions instead. +#: Neighbor tables keyed by the connectivity's `offset_tag`, which is how the IR names it. OffsetProvider: TypeAlias = Mapping[Tag, OffsetProviderElem] OffsetProviderType: TypeAlias = Mapping[Tag, OffsetProviderTypeElem] +#: An offset provider as users write it: keyed by `NeighborConnectivity` declarations (or, at the +#: IR level, by tags). The entry points of a program normalize it to an `OffsetProvider` with +#: `as_tag_keyed_offset_provider`, so everything below them sees tags only. +#: NOTE: `Any` keys, since `Mapping` is invariant in its key type: a tag-keyed provider would not +#: be a `Mapping[type[NeighborConnectivity] | Tag, ...]`. Keys are checked at runtime instead. +OffsetProviderLike: TypeAlias = Mapping[Any, OffsetProviderElem] +OffsetProviderTypeLike: TypeAlias = Mapping[Any, OffsetProviderTypeElem] def is_offset_provider(obj: Any) -> TypeGuard[OffsetProvider]: @@ -1450,13 +1534,59 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid """ # TODO(havogt): Once we have a custom class for `OffsetProvider`, we can absorb this functionality into it. if offset_tag not in offset_provider: - raise KeyError(f"Offset '{offset_tag}' not found in offset provider.") - return offset_provider[offset_tag] # TODO return a valid dimension + raise KeyError( + f"Connectivity '{offset_tag}' not found in the offset provider, which has" + f" {sorted(map(str, offset_provider))}. Offset providers are keyed by" + " 'NeighborConnectivity' declarations, e.g. '{V2E: v2e_table}'." + ) + return offset_provider[offset_tag] get_offset_type: Callable[[OffsetProviderType, str], OffsetProviderTypeElem] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and OffsetProviderType overlap +def connectivity_key_over( + offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension | Tag +) -> str: + """ + The key of a bound connectivity whose local dimension is `local_dim` (a dimension or its tag). + + Neighbor reductions and sparse fields know only their local dimension, and use its table + for the neighbor count and the skip values. That is the table keyed by the local dimension's + tag, i.e. its owner's, if bound. Otherwise it is one of the connectivities *sharing* the local + dimension (see `NeighborConnectivity`), each keyed by its own tag; the smallest key is taken, + so the choice does not depend on the order of the provider. Connectivities sharing a local + dimension have the same neighbor structure (see `check_offset_provider`), so which one does + not matter. + + Raises: + KeyError: If no bound connectivity has `local_dim` as its local dimension. + """ + local_tag = local_dim if isinstance(local_dim, str) else local_dim.tag + if local_tag in offset_provider: + return local_tag + candidates = [ + key + for key, connectivity in offset_provider.items() + if (neighbor_dim := _neighbor_dim_of(connectivity)) is not None + and neighbor_dim.tag == local_tag + ] + if not candidates: + raise KeyError( + f"No connectivity over the local dimension '{local_tag}' is bound in the offset" + f" provider, which has {sorted(map(str, offset_provider))}." + ) + return min(candidates) + + +def _neighbor_dim_of(connectivity: Any) -> Optional[Dimension]: + if isinstance(connectivity, NeighborConnectivityType): + return connectivity.neighbor_dim + if is_neighbor_table(connectivity): + return connectivity.domain.dims[1] + return None + + def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: str) -> bool: """Determine if offset provider has an element for the given offset tag.""" try: @@ -1466,7 +1596,9 @@ def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: return True -def hash_offset_provider_items_by_id(offset_provider: OffsetProvider) -> int: +def hash_offset_provider_items_by_id( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, +) -> int: """ Compute hash of an offset provider on the tuples of key and value id. @@ -1557,8 +1689,8 @@ def inverse_image(self, image_range: UnitRange | NamedRange) -> Sequence[NamedRa def premap( self, - index_field: Connectivity | fbuiltins.FieldOffset, - *args: Connectivity | fbuiltins.FieldOffset, + index_field: Connectivity | type[NeighborConnectivity], + *args: Connectivity | type[NeighborConnectivity], ) -> Connectivity: raise NotImplementedError() @@ -1627,8 +1759,8 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: >>> class I(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... >>> class J(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... >>> class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... - >>> class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... - >>> class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... + >>> class E2V(LocalDimensionIndex): ... + >>> class E2C(LocalDimensionIndex): ... >>> promote_dims([J, K], [I, K]) == [I, J, K] True >>> promote_dims([K, J], [I, K]) @@ -1731,6 +1863,8 @@ def __getitem__(cls, base: Dimension) -> Dimension: ) if not isinstance(base, DimensionMeta): raise TypeError(f"'Staggered' expects a dimension, got '{base!r}'.") + if base.kind is DimensionKind.LOCAL: + raise TypeError(f"'{base.__qualname__}' is a local dimension and cannot be staggered.") if is_staggered(base): raise TypeError( f"'{base.__qualname__}' is already staggered; a dimension cannot be staggered twice." @@ -1791,24 +1925,6 @@ def __init_subclass__(cls, /, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) -class ConstListDim(DimensionIndex, kind=DimensionKind.LOCAL): - """ - The local dimension of a list whose length is known at compile time (`make_const_list`). - - Declared here, once, because it must be a *single* class. It used to be built - independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless - while dimensions compared by `(name, kind)` -- the two instances were equal. Under - nominal identity (ADR 0028) two declarations would be two different dimensions, and the - `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s - built by embedded execution. - - TODO: becomes an owner-less local dimension with an explicit size, generalising this from - length 1 to length *n*, once local dimensions know their connectivity. - """ - - __slots__ = () - - def _reduce_staggered(cls: StaggeredMeta) -> Any: """ Pickle a staggered dimension through its base, falling back to by-reference. @@ -1878,3 +1994,550 @@ def connectivity_for_cartesian_shift(dim: Dimension, offset: int | float) -> Car else: assert staggered_offset == 0 return CartesianConnectivity(dim, int(integral_offset), codomain=dim) + + +class LocalDimensionIndex(DimensionIndex): + """ + A local dimension: the axis that runs over the neighbors of one element. + + A local dimension is declared either inside a `NeighborConnectivity`, as its nested `Local` + class, or on its own for a local axis that indexes no table (`owner is None`), such as the + coefficients of a fixed-size stencil: + + >>> class LsqCoeff(LocalDimensionIndex, size=3): ... + >>> LsqCoeff.kind, LsqCoeff.owner, LsqCoeff.max_neighbors + (, None, 3) + + Neighbor counts are optional. A declared count is a constraint the bound table has to + satisfy (see `check_neighbor_table`); an undeclared one is taken from the table. + """ + + __slots__ = () + + kind: ClassVar[DimensionKind] = DimensionKind.LOCAL + _local_dimension_root: ClassVar[bool] = True + + #: The connectivity this dimension is the local axis of, or `None` if it indexes no table. + #: Set by `NeighborConnectivity` when the connectivity is declared. + owner: ClassVar[Optional[type[NeighborConnectivity]]] = None + #: Number of entries per element, i.e. the table's second extent, if declared. + max_neighbors: ClassVar[Optional[int]] = None + #: Least number of *valid* neighbors of any element, if declared. Fewer than + #: `max_neighbors` means the table pads with skip values. + min_neighbors: ClassVar[Optional[int]] = None + #: The `size=` of this declaration, kept apart from the counts an owner writes below. + declared_size: ClassVar[Optional[int]] = None + + def __init_subclass__( + cls, + /, + *, + size: Optional[int] = None, + kind: Optional[DimensionKind] = None, + **kwargs: Any, + ) -> None: + if kind is not None and kind is not DimensionKind.LOCAL: + raise TypeError( + f"'{cls.__qualname__}' is a local dimension and cannot have kind '{kind}'." + ) + super().__init_subclass__(**kwargs) + # NOTE: reset rather than inherited: a subclass of an owned local dimension is a + # different dimension, and does not index its parent's table. + cls.owner = None + cls.declared_size = _check_neighbor_count(cls, "size", size) + cls.max_neighbors = cls.min_neighbors = cls.declared_size + + +def _check_neighbor_count(cls: type, name: str, count: Optional[int]) -> Optional[int]: + if count is None: + return None + if not isinstance(count, numbers.Integral) or isinstance(count, bool): + raise TypeError(f"'{cls.__qualname__}': '{name}' must be an integer, got '{count!r}'.") + if count < 0: + raise ValueError(f"'{cls.__qualname__}': '{name}' must be non-negative, got {count}.") + return int(count) + + +class ConstList(LocalDimensionIndex, size=1): + """ + The local dimension of a list of one repeated value (`make_const_list`). + + An owner-less local dimension of size 1: the value is broadcast against the neighbor lists it + is combined with, and a materialized constant list has extent 1 along it. It indexes no table, + so it is never in an offset provider. + + Declared here, once: it used to be built independently in `iterator/embedded.py` and in the + DaCe lowering, which only worked while dimensions compared by `(name, kind)`. + """ + + __slots__ = () + + +class ConnectivityMeta(type): + """ + Metaclass of `NeighborConnectivity` declarations. + + A connectivity declaration is a class that is never instantiated. It is written in DSL code + (`a(V2E)`, `a(V2E[0])`), and it is what the neighbor table bound at call time must match. + """ + + # NOTE: `Local` is deliberately *not* annotated here, nor on `NeighborConnectivity`: an + # annotated `Local` makes every declaration's nested class a *variable* for the checkers, so + # `Field[Dims[V, V2E.Local]]` is rejected (pyright) or "not valid as a type" (mypy, for the + # assigned form). Library code reads it through `local_dimension_of`. + origin: Dimension + codomain: Dimension + + @property + def tag(cls) -> Tag: + """The connectivity's identity: its qualified Python name.""" + return f"{cls.__module__}.{cls.__qualname__}" + + @property + def offset_tag(cls) -> Tag: + """ + The name of the connectivity in the IR, and its key in a normalized offset provider. + + The tag of its local dimension, if it declares it: shifts, neighbor reductions and sparse + arguments then all find the table under one string. A connectivity that *shares* another + one's local dimension (a flattened sparse pattern, e.g. cell-to-cell-edge indexing the + same neighbor axis as cell-to-edge) is named by its own tag, since the local dimension's + tag already names its owner's table. + """ + local = cls._local() + return local.tag if local.owner is cls else cls.tag + + def _local(cls) -> type[LocalDimensionIndex]: + if (local := cls.__dict__.get("Local")) is None: + raise TypeError( + f"'{cls.__qualname__}' is not a connectivity declaration; declare one by" + " subclassing 'NeighborConnectivity[Origin, Codomain]'." + ) + return cast(type[LocalDimensionIndex], local) + + def __call__(cls, *args: Any, **kwargs: Any) -> NoReturn: + raise TypeError( + f"'{cls.__qualname__}' is a connectivity declaration and cannot be instantiated;" + " bind a neighbor table to it through the offset provider." + ) + + @overload + def __getitem__(cls, item: int) -> Connectivity: ... + @overload + def __getitem__(cls, item: Any) -> Any: ... + def __getitem__(cls, item: Any) -> Any: + # NOTE: `numbers.Integral`, not `int`, so `V2E[np.int32(1)]` does not fall through to + # type-parameter subscription; `bool` is excluded so `V2E[True]` is an error. + if isinstance(item, numbers.Integral) and not isinstance(item, bool): + return cls.bound_table()[cls._local()(int(item))] + if "Local" in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}[{item!r}]': a connectivity is indexed by an integer" + " neighbor position." + ) + # A metaclass `__getitem__` shadows `__class_getitem__`, so type-parameter + # subscription (`NeighborConnectivity[V, E]`) has to be forwarded explicitly. + return cast(Any, cls).__class_getitem__(item) + + def __repr__(cls) -> str: + return cls.tag + + def __str__(cls) -> str: + return cls.__qualname__ + + def __gt_type__(cls) -> Any: + """ + The type of the connectivity in DSL code: an offset from `Codomain` to `(Origin, Local)`. + + Its tag is `offset_tag`, which is how the IR names the connectivity and how the offset + provider is keyed once normalized (see `as_tag_keyed_offset_provider`). + """ + from gt4py.next.type_system import type_specifications as ts + + local = cls._local() + return ts.OffsetType(source=cls.codomain, target=(cls.origin, local), tag=cls.offset_tag) + + def bound_table(cls) -> NeighborTable: + """The neighbor table bound to this connectivity in the current embedded execution.""" + from gt4py.next import embedded + + offset_provider = embedded.context.get_offset_provider(None) + if offset_provider is None: + raise RuntimeError( + f"'{cls.__qualname__}' can only be resolved to a table during embedded execution." + ) + table = get_offset(offset_provider, cls.offset_tag) + if not is_neighbor_table(table): + raise TypeError( + f"'{cls.__qualname__}' is bound to '{table}', which is not a neighbor table." + ) + return table + + +class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + """ + Declare a neighbor connectivity: for each `Origin` element, a list of `Codomain` neighbors. + + The declaration names the connectivity's local dimension -- its nested `Local` class -- + and optionally its neighbor counts. It holds no data: the neighbor table is bound at call + time through the offset provider. `check_neighbor_table` checks a table against the + declaration. + + Examples: + >>> class Vertex(DimensionIndex): ... + >>> class Edge(DimensionIndex): ... + >>> class V2E(NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): + ... class Local(LocalDimensionIndex): ... + >>> V2E.origin is Vertex, V2E.codomain is Edge + (True, True) + >>> V2E.Local.owner is V2E, V2E.Local.max_neighbors, V2E.Local.min_neighbors + (True, 6, 5) + """ + + # NOTE: `Local` is not annotated (see `ConnectivityMeta`); every subclass declares it, as a + # nested class or as `Local: TypeAlias = `. + origin: ClassVar[Dimension] + codomain: ClassVar[Dimension] + + def __init_subclass__( + cls, + /, + *, + max_neighbors: Optional[int] = None, + min_neighbors: Optional[int] = None, + **kwargs: Any, + ) -> None: + super().__init_subclass__(**kwargs) + name = cls.__qualname__ + if "" in name: + raise TypeError( + f"'{name}' must be declared at module level: a connectivity is referenced from" + " the IR by its qualified name, which has to be importable." + ) + params = [ + xtyping.get_args(base) + for base in cls.__dict__.get("__orig_bases__", ()) + if xtyping.get_origin(base) is NeighborConnectivity + ] + if len(params) != 1 or len(params[0]) != 2: + raise TypeError( + f"'{name}' must derive from 'NeighborConnectivity[Origin, Codomain]' directly," + " with both dimensions given." + ) + origin, codomain = params[0] + for role, dim in (("Origin", origin), ("Codomain", codomain)): + if not isinstance(dim, DimensionMeta) or dim.kind is DimensionKind.LOCAL: + raise TypeError(f"'{name}': '{role}' must be a non-local dimension, got '{dim}'.") + + local = cls.__dict__.get("Local") + if not (isinstance(local, DimensionMeta) and issubclass(local, LocalDimensionIndex)): + raise TypeError( + f"'{name}' must declare its local dimension, either as a nested class" + " ('class Local(LocalDimensionIndex): ...') or by adopting one" + " ('Local: TypeAlias = SomeLocalDim')." + ) + if local is ConstList: + raise TypeError( + f"'{name}' cannot adopt '{ConstList.__qualname__}': it is the local dimension" + " of 'make_const_list' results and belongs to no connectivity." + ) + max_neighbors = _check_neighbor_count(cls, "max_neighbors", max_neighbors) + min_neighbors = _check_neighbor_count(cls, "min_neighbors", min_neighbors) + # NOTE: a declaration whose tag is the owner's is a *redefinition* of it (a re-run + # notebook cell), not a second connectivity sharing the local dimension, so it takes + # ownership over again. Ownership of a local dimension that is adopted, rather than + # nested, otherwise goes to whoever declares first. + if local.owner is not None and local.owner.tag != cls.tag: + # Sharing another connectivity's local dimension: the neighbor structure is the + # owner's, including its counts, and the sharing connectivity is named by its own tag. + owner_name = local.owner.__qualname__ + if origin is not local.owner.origin: + raise TypeError( + f"'{name}' cannot share the local dimension of '{owner_name}': it has origin" + f" '{origin}', but the neighbors of '{owner_name}' are those of" + f" '{local.owner.origin}'." + ) + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("min_neighbors", min_neighbors), + ): + if count is not None and count != getattr(local, count_name): + raise TypeError( + f"'{name}': '{count_name}={count}' contradicts the local dimension it" + f" shares with '{owner_name}', which declares" + f" {count_name}={getattr(local, count_name)}." + ) + cls.origin, cls.codomain = origin, codomain + return + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("min_neighbors", min_neighbors), + ): + # NOTE: against the local dimension's own `size=`, not against counts a previous + # owner wrote: a redefinition must be checked against what its `Local` declares. + if ( + count is not None + and local.declared_size is not None + and count != local.declared_size + ): + raise TypeError( + f"'{name}': '{count_name}={count}' contradicts the size declared by" + f" '{local.__qualname__}' ({local.declared_size})." + ) + max_neighbors = max_neighbors if max_neighbors is not None else local.declared_size + min_neighbors = min_neighbors if min_neighbors is not None else local.declared_size + if ( + max_neighbors is not None + and min_neighbors is not None + and min_neighbors > max_neighbors + ): + raise TypeError( + f"'{name}': 'min_neighbors' ({min_neighbors}) exceeds 'max_neighbors'" + f" ({max_neighbors})." + ) + + cls.origin, cls.codomain = origin, codomain + local.owner = cls + local.max_neighbors, local.min_neighbors = max_neighbors, min_neighbors + + +def local_dimension_of(connectivity: type[NeighborConnectivity]) -> type[LocalDimensionIndex]: + """ + The local dimension a connectivity declares, adopts or shares. + + Library code reads `V2E.Local` through this accessor: the attribute is intentionally not + annotated, so that a declaration's `Local` stays a *type* for the type checkers (see + `ConnectivityMeta`). + + Raises: + TypeError: If `connectivity` declares no local dimension. + """ + return cast(ConnectivityMeta, connectivity)._local() + + +def check_neighbor_table( + connectivity: type[NeighborConnectivity], + table: NeighborTable | NeighborConnectivityType, +) -> None: + """ + Check that a neighbor table matches the connectivity declaration it is bound to. + + Skip values are checked on the table's *type*: a table with a `skip_value` counts as + having skip values whether or not any entry uses it. + + Args: + connectivity: The declaration. + table: The bound table, or its type (which is all an ahead-of-time compilation has). + + Raises: + ValueError: On the first mismatch, naming the connectivity and the mismatch. + """ + table_type = table if isinstance(table, NeighborConnectivityType) else table.__gt_type__() + name = connectivity.__qualname__ + local = local_dimension_of(connectivity) + + def fail(reason: str) -> NoReturn: + raise ValueError(f"The table bound to '{name}' does not match its declaration: {reason}.") + + if not isinstance(table_type, NeighborConnectivityType): + fail(f"expected a neighbor table, got '{table_type}'") + + def redefined(found: Sequence[Dimension], expected: Sequence[Dimension]) -> str: + if any(f is not e and f.tag == e.tag for f, e in zip(found, expected)): + return ( + " (a dimension of the same name but a different class: was the declaration" + " redefined, e.g. by re-running a notebook cell?)" + ) + return "" + + expected_domain = (connectivity.origin, local) + if tuple(table_type.domain) != expected_domain: + fail( + f"its domain is '({', '.join(map(str, table_type.domain))})'," + f" expected '({', '.join(map(str, expected_domain))})'" + + redefined(table_type.domain, expected_domain) + ) + if table_type.codomain is not connectivity.codomain: + fail( + f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'" + + redefined((table_type.codomain,), (connectivity.codomain,)) + ) + if not np.issubdtype(table_type.dtype.scalar_type, np.integer): + fail(f"its dtype '{table_type.dtype}' is not integral") + if local.max_neighbors is not None and table_type.max_neighbors != local.max_neighbors: + fail( + f"it has {table_type.max_neighbors} neighbors per element," + f" expected max_neighbors={local.max_neighbors}" + ) + if local.min_neighbors is not None: + max_neighbors = table_type.max_neighbors + if local.min_neighbors > max_neighbors: + fail( + f"min_neighbors={local.min_neighbors} exceeds its {max_neighbors} neighbors" + " per element" + ) + if local.min_neighbors < max_neighbors and not table_type.has_skip_values: + fail( + f"min_neighbors={local.min_neighbors} < {max_neighbors} requires a skip value," + " but the table has none" + ) + if local.min_neighbors == max_neighbors and table_type.has_skip_values: + fail( + f"min_neighbors == max_neighbors == {max_neighbors} means every element has all" + f" its neighbors, but the table has skip value {table_type.skip_value}" + ) + + +@overload +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderLike, *, strict: bool = True +) -> OffsetProvider: ... +@overload +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderTypeLike, *, strict: bool = True +) -> OffsetProviderType: ... +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, *, strict: bool = True +) -> OffsetProvider | OffsetProviderType: + """ + Key an offset provider by tags, the form the IR and the backends use. + + A `NeighborConnectivity` key becomes its `offset_tag`: its local dimension's tag, or its own + for a connectivity sharing another one's local dimension. A string key is taken to be + such a tag already, and is rejected if it cannot be one: a tag is a qualified name, so a bare + name such as `"V2E"` is the removed `FieldOffset` spelling. + + Called on every program call, so it does not check tables against their declarations; see + `check_offset_provider`. + + Args: + offset_provider: The provider to normalize. + strict: Whether to reject string keys that cannot be tags. Internal hooks that are handed + hand-written providers, such as `embedded.context.update`, pass `False`. + """ + if not any(isinstance(key, ConnectivityMeta) for key in offset_provider): + if strict: + _check_tag_keys(offset_provider) + return offset_provider + result: dict[Tag, Any] = {} + for key, value in offset_provider.items(): + tag = key.offset_tag if isinstance(key, ConnectivityMeta) else key + if tag in result: + raise ValueError(f"The offset provider binds '{tag}' twice.") + result[tag] = value + if strict: + _check_tag_keys(result) + return result + + +def _check_tag_keys(offset_provider: Mapping[Any, Any]) -> None: + for key in offset_provider: + if not isinstance(key, str) or "." not in key: + raise TypeError( + f"Invalid offset-provider key {key!r}: offset providers are keyed by" + " 'NeighborConnectivity' declarations, e.g. '{V2E: v2e_table}'. A bare name is the" + " spelling of the removed 'FieldOffset' (see ADR 0029)." + ) + + +#: Offset providers already checked, by the hash of their `(key, id(table))` items. Bounded, and +#: not authoritative: like the compiled-program cache (which keys on the same hash), it can in +#: principle skip a check when a freed table is replaced at the same address. See +#: `check_offset_provider`. +_CHECKED_OFFSET_PROVIDERS: Final[collections.OrderedDict[int, None]] = collections.OrderedDict() +_CHECKED_OFFSET_PROVIDERS_MAX: Final = 256 + + +def check_offset_provider( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, *, deep: bool = False +) -> None: + """ + Check every table of an offset provider against its connectivity declaration. + + A key that does not name a declared connectivity -- e.g. a tag used only by hand-written IR -- + has no declaration to be checked against and is skipped. Providers are remembered by the + identity of their tables, so repeated calls with the same tables cost one hash. + + Args: + offset_provider: The provider, keyed by declarations or by tags. + deep: Also compare the skip-value positions of tables over one shared local dimension, + which reads the tables. The compile path does; the call path does not. + + Raises: + ValueError: If a table does not match its declaration, see `check_neighbor_table`. + """ + if (seen := hash_offset_provider_items_by_id(offset_provider)) in _CHECKED_OFFSET_PROVIDERS: + return + for key, table in offset_provider.items(): + declaration: Any = key + if isinstance(key, str): + if (declaration := _import_qualified_name_or_none(key)) is None: + continue + if isinstance(declaration, DimensionMeta): + # the local dimension's tag names its owner's table + declaration = getattr(declaration, "owner", None) + if isinstance(declaration, ConnectivityMeta): + if declaration.offset_tag not in (key, getattr(key, "offset_tag", None)): + # e.g. `{V2E.tag: table}`: the connectivity's own tag, which is the IR name only + # of a connectivity *sharing* a local dimension + raise ValueError( + f"Invalid offset-provider key '{key}': it names the connectivity" + f" '{declaration.__qualname__}', whose key is the declaration itself" + f" ('{{{declaration.__qualname__}: table}}')." + ) + check_neighbor_table(cast(type[NeighborConnectivity], declaration), table) + _check_shared_local_dimensions(offset_provider, deep=deep) + _CHECKED_OFFSET_PROVIDERS[seen] = None + while len(_CHECKED_OFFSET_PROVIDERS) > _CHECKED_OFFSET_PROVIDERS_MAX: + _CHECKED_OFFSET_PROVIDERS.popitem(last=False) + + +def _check_shared_local_dimensions( + offset_provider: OffsetProviderLike | OffsetProviderTypeLike, *, deep: bool = False +) -> None: + """ + Check that the tables over one local dimension have the same neighbor structure. + + Reductions and sparse fields take the neighbor count and the skip values of a local dimension + from any one table over it (see `connectivity_key_over`), which is only sound if all of them + agree: the same number of neighbors, and a skip value at the same positions. + """ + by_local_dim: dict[Tag, list[tuple[Any, Any]]] = collections.defaultdict(list) + for key, table in offset_provider.items(): + if (neighbor_dim := _neighbor_dim_of(table)) is not None: + by_local_dim[neighbor_dim.tag].append((key, table)) + for local_tag, tables in by_local_dim.items(): + (first_key, first), *others = tables + first_type = first if isinstance(first, NeighborConnectivityType) else first.__gt_type__() + for key, table in others: + table_type = ( + table if isinstance(table, NeighborConnectivityType) else table.__gt_type__() + ) + same_structure = (table_type.max_neighbors, table_type.has_skip_values) == ( + first_type.max_neighbors, + first_type.has_skip_values, + ) + if ( + deep + and same_structure + and first_type.has_skip_values + and is_neighbor_table(first) + and is_neighbor_table(table) + ): + # NOTE: compared where the tables live, without copying device arrays to the host. + xp = first.array_ns # type: ignore[attr-defined] # all tables are NdArrayFields + same_structure = first.ndarray.shape == table.ndarray.shape and bool( + xp.all( + (first.ndarray == first_type.skip_value) + == (xp.asarray(table.ndarray) == table_type.skip_value) + ) + ) + if not same_structure: + raise ValueError( + f"'{key}' and '{first_key}' are bound to tables over the same local dimension" + f" '{local_tag}' with a different neighbor structure: connectivities sharing a" + " local dimension must have the same number of neighbors, and skip values at" + " the same positions." + ) diff --git a/src/gt4py/next/constructors.py b/src/gt4py/next/constructors.py index f1434bc954..ba6c1e5cab 100644 --- a/src/gt4py/next/constructors.py +++ b/src/gt4py/next/constructors.py @@ -653,7 +653,7 @@ def as_connectivity( >>> from gt4py import next as gtx >>> class Vertex(gtx.DimensionIndex): ... >>> class Edge(gtx.DimensionIndex): ... - >>> class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... + >>> class V2EDim(gtx.LocalDimensionIndex): ... >>> data = np.array([[0, 1], [1, 2], [2, 0]]) >>> conn = gtx.as_connectivity([Vertex, V2EDim], Edge, data) >>> conn.ndarray diff --git a/src/gt4py/next/embedded/context.py b/src/gt4py/next/embedded/context.py index 8183a3292c..6ca2544ce9 100644 --- a/src/gt4py/next/embedded/context.py +++ b/src/gt4py/next/embedded/context.py @@ -73,7 +73,7 @@ def get_offset_provider(default: _T = _NO_DEFAULT_SENTINEL) -> common.OffsetProv def update( *, closure_column_range: common.NamedRange | eve.NothingType = eve.NOTHING, - offset_provider: common.OffsetProvider | eve.NothingType = eve.NOTHING, + offset_provider: common.OffsetProviderLike | eve.NothingType = eve.NOTHING, ) -> Generator[None, None, None]: """Context handler updating the current embedded context with the provided values.""" @@ -83,7 +83,9 @@ def update( closure_token = gtx_embedded.context._closure_column_range.set(closure_column_range) if offset_provider is not eve.NOTHING: assert not isinstance(offset_provider, eve.NothingType) - offset_provider_token = gtx_embedded.context._offset_provider.set(offset_provider) + offset_provider_token = gtx_embedded.context._offset_provider.set( + common.as_tag_keyed_offset_provider(offset_provider, strict=False) + ) try: yield None diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index 0e3aaeab8b..114442100f 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -239,7 +239,7 @@ def as_scalar(self) -> core_defs.ScalarT: def premap( self: NdArrayField, - *connectivities: common.Connectivity | fbuiltins.FieldOffset, + *connectivities: common.Connectivity | type[common.NeighborConnectivity], ) -> NdArrayField: """ Rearrange the field content using the provided connectivities (index mappings). @@ -314,10 +314,9 @@ def premap( codomains_counter: collections.Counter[common.Dimension] = collections.Counter() for connectivity in connectivities: - # For neighbor reductions, a FieldOffset is passed instead of an actual Connectivity - if not isinstance(connectivity, common.Connectivity): - assert isinstance(connectivity, fbuiltins.FieldOffset) - connectivity = connectivity.as_connectivity_field() + # For neighbor reductions, a connectivity declaration is passed instead of a table + if isinstance(connectivity, common.ConnectivityMeta): + connectivity = connectivity.bound_table() assert isinstance(connectivity, common.Connectivity) # Current implementation relies on skip_value == -1: @@ -366,8 +365,8 @@ def premap( def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: return functools.reduce( lambda field, current_index_field: field.premap(current_index_field), @@ -941,15 +940,12 @@ def _concat_where( NdArrayField.register_builtin_func(experimental.concat_where, _concat_where) # type: ignore[arg-type] -def _as_offset(offset: fbuiltins.FieldOffset, offset_field: NdArrayField) -> common.Connectivity: - if not fbuiltins.is_cartesian_offset(offset): - target_dims = ", ".join(d.__qualname__ for d in offset.target) # for the diagnostic - raise ValueError( - f"'as_offset' is only supported for Cartesian offsets " - f"(single target dimension equal to source dimension); " - f"got source '{offset.source.__qualname__}' and target ({target_dims})." - ) - source_dim = offset.source +def _as_offset(source_dim: common.Dimension, offset_field: NdArrayField) -> common.Connectivity: + if ( + not isinstance(source_dim, common.DimensionMeta) + or source_dim.kind is common.DimensionKind.LOCAL + ): + raise ValueError(f"'as_offset' shifts along a non-local dimension, got '{source_dim}'.") coords = _identity_index_array( offset_field.domain, source_dim, offset_field.array_ns, dtype=fbuiltins.IndexType ) @@ -981,8 +977,8 @@ def _builtin_op( current_offset_provider = embedded_context.get_offset_provider(None) assert current_offset_provider is not None offset_definition = common.get_offset( - current_offset_provider, axis.tag - ) # assumes offset and local dimension have same name + current_offset_provider, common.connectivity_key_over(current_offset_provider, axis) + ) assert common.is_neighbor_table(offset_definition) new_domain = common.Domain(*[nr for nr in field.domain if nr.dim != axis]) diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index e9daedb8a9..d409428b24 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -160,9 +160,9 @@ def _make_compiled_programs_pool( def compile( self, - offset_provider: common.OffsetProviderType - | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + offset_provider: common.OffsetProviderTypeLike + | common.OffsetProviderLike + | list[common.OffsetProviderTypeLike | common.OffsetProviderLike] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -199,6 +199,7 @@ def compile( ) if not isinstance(offset_provider, list): offset_provider = [offset_provider] # type: ignore[list-item] # cleanup offset_provider vs offset_provider_type + offset_provider = [common.as_tag_keyed_offset_provider(op) for op in offset_provider] assert all( common.is_offset_provider(op) or common.is_offset_provider_type(op) @@ -376,12 +377,13 @@ def with_bound_args(self, **kwargs: Any) -> ProgramWithBoundArgs: def __call__( self, *args: Any, - offset_provider: common.OffsetProvider | None = None, + offset_provider: common.OffsetProviderLike | None = None, enable_jit: bool | None = None, **kwargs: Any, ) -> None: - if offset_provider is None: - offset_provider = {} + offset_provider = ( + {} if offset_provider is None else common.as_tag_keyed_offset_provider(offset_provider) + ) enable_jit = self.compilation_options.enable_jit if enable_jit is None else enable_jit with program_call_context( @@ -412,6 +414,7 @@ def __call__( stacklevel=2, ) + common.check_offset_provider(offset_provider) with next_embedded.context.update(offset_provider=offset_provider): with embedded_program_call_context(self, args, offset_provider, kwargs): self.definition_stage.definition(*args, **kwargs) @@ -431,7 +434,7 @@ class ProgramWithBoundArgs(Program): @override def __call__( - self, *args: Any, offset_provider: common.OffsetProvider | None = None, **kwargs: Any + self, *args: Any, offset_provider: common.OffsetProviderLike | None = None, **kwargs: Any ) -> None: if offset_provider is None: offset_provider = {} @@ -488,9 +491,9 @@ def __call__( @override def compile( self, - offset_provider: common.OffsetProviderType - | common.OffsetProvider - | list[common.OffsetProviderType | common.OffsetProvider] + offset_provider: common.OffsetProviderTypeLike + | common.OffsetProviderLike + | list[common.OffsetProviderTypeLike | common.OffsetProviderLike] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -655,7 +658,9 @@ def __gt_closure_vars__(self) -> dict[str, Any]: def __call__(self, *args: Any, enable_jit: bool | None = None, **kwargs: Any) -> Any: if not next_embedded.context.within_valid_context() and self.backend is not None: # non embedded execution - offset_provider = {**kwargs.pop("offset_provider", {})} + offset_provider = { + **common.as_tag_keyed_offset_provider(kwargs.pop("offset_provider", {})) + } if "out" not in kwargs: raise errors.MissingArgumentError(None, "out", True) out = kwargs.pop("out") @@ -677,7 +682,10 @@ def __call__(self, *args: Any, enable_jit: bool | None = None, **kwargs: Any) -> else: if not next_embedded.context.within_valid_context(): # field_operator as program - kwargs["offset_provider"] = {**kwargs.pop("offset_provider", {})} + kwargs["offset_provider"] = { + **common.as_tag_keyed_offset_provider(kwargs.pop("offset_provider", {})) + } + common.check_offset_provider(kwargs["offset_provider"]) attributes = ( self.definition_stage.attributes if self.definition_stage @@ -721,6 +729,11 @@ class FieldOperatorFromFoast(FieldOperator): @override def __call__(self, *args: Any, **kwargs: Any) -> Any: assert self.backend is not None + if "offset_provider" in kwargs: + kwargs["offset_provider"] = common.as_tag_keyed_offset_provider( + kwargs["offset_provider"] + ) + common.check_offset_provider(kwargs["offset_provider"]) compiled_fo = self.backend.compile( self.foast_stage, arguments.CompileTimeArgs.from_concrete(*args, **kwargs) ) diff --git a/src/gt4py/next/ffront/experimental.py b/src/gt4py/next/ffront/experimental.py index b547d67663..b1824b61f2 100644 --- a/src/gt4py/next/ffront/experimental.py +++ b/src/gt4py/next/ffront/experimental.py @@ -10,11 +10,16 @@ from gt4py._core import definitions as core_defs from gt4py.next import common, named_collections -from gt4py.next.ffront.fbuiltins import BuiltInFunction, FieldOffset, WhereBuiltinFunction +from gt4py.next.ffront.fbuiltins import BuiltInFunction, WhereBuiltinFunction @BuiltInFunction -def as_offset(offset: FieldOffset, field: common.Field, /) -> common.Connectivity: +def as_offset(dim: common.Dimension, field: common.Field, /) -> common.Connectivity: + """ + Shift along `dim` by the per-point amounts in the integer `field`. + + `a(as_offset(KDim, k_offsets))` reads `a` at `k + k_offsets[k]` in `KDim`. + """ raise NotImplementedError() diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 7474cd1406..cf7d4a2ff4 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -7,7 +7,6 @@ # SPDX-License-Identifier: BSD-3-Clause import dataclasses -import functools import inspect import math import operator @@ -35,7 +34,6 @@ from gt4py._core import definitions as core_defs from gt4py.next import common, named_collections from gt4py.next.common import Dimension, Field # noqa: F401 [unused-import] for TYPE_BUILTINS -from gt4py.next.iterator import runtime from gt4py.next.type_system import type_specifications as ts @@ -129,8 +127,6 @@ def _type_conversion_helper(t: type) -> type[ts.TypeSpec] | tuple[type[ts.TypeSp return ts.FieldType elif t is common.Dimension: return ts.DimensionType - elif t is FieldOffset: - return ts.OffsetType elif t is common.Connectivity: return ts.OffsetType elif t is core_defs.ScalarT: @@ -472,70 +468,3 @@ def impl( assert (diff := actual_export - should_export) == set(), ( f"Symbol(s) exported but not defined in 'fbuiltins': {diff}" ) - - -# TODO(tehrengruber): FieldOffset and runtime.Offset are not an exact conceptual -# match. Revisit if we want to continue subclassing here. If we split -# them also check whether Dimension should continue to be the shared or define -# guidelines for decision. -@dataclasses.dataclass(frozen=True) -class FieldOffset(runtime.Offset): - #: The tag, i.e. the offset-provider key. - value: str - source: common.Dimension - target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] - - @functools.cached_property - def _cache(self) -> dict: - return {} - - def __post_init__(self) -> None: - if len(self.target) == 2 and self.target[1].kind != common.DimensionKind.LOCAL: - raise ValueError("Second dimension in offset must be a local dimension.") - - def __gt_type__(self) -> ts.OffsetType: - return ts.OffsetType(source=self.source, target=self.target, tag=self.value) - - def __getitem__(self, offset: int) -> common.Connectivity: - """Serve as a connectivity factory.""" - from gt4py.next import embedded # avoid circular import - - assert isinstance(self.value, str) - current_offset_provider = embedded.context.get_offset_provider(None) - assert current_offset_provider is not None - offset_definition = common.get_offset(current_offset_provider, self.value) - - assert common.is_neighbor_table(offset_definition) - named_index = (self.target[-1])(offset) - connectivity = offset_definition[named_index] - - return connectivity - - def as_connectivity_field(self) -> common.Connectivity: - """Convert to connectivity field using the offset providers in current embedded execution context.""" - from gt4py.next import embedded # avoid circular import - - assert isinstance(self.value, str) - current_offset_provider = embedded.context.get_offset_provider(None) - assert current_offset_provider is not None - offset_definition = common.get_offset(current_offset_provider, self.value) - - cache_key = id(offset_definition) - if (connectivity := self._cache.get(cache_key, None)) is None: - if isinstance(offset_definition, common.Connectivity): - connectivity = offset_definition - else: - raise NotImplementedError() - - self._cache[cache_key] = connectivity - - return connectivity - - -def is_cartesian_offset(offset: FieldOffset | ts.OffsetType) -> bool: - return ( - len(offset.target) == 1 - and offset.source == offset.target[0] - and offset.source.kind == offset.target[0].kind - and offset.target[0].kind != common.DimensionKind.LOCAL - ) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 3d09d7ecb4..9aa62ecae7 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -434,11 +434,20 @@ def visit_Symbol( def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> foast.Attribute: new_value = self.visit(node.value, **kwargs) + match new_value.type: + # `V2E.Local`: the local dimension of a connectivity declaration, which is the last + # target of the offset it is typed as. + case ts.OffsetType(target=(_, local)) if node.attr == "Local": + attr_type: ts.TypeSpec = ts.DimensionType(dim=local) + case _: + try: + attr_type = getattr(new_value.type, node.attr) + except AttributeError: + raise errors.DSLError( + node.location, f"'{new_value.type}' has no attribute '{node.attr}'." + ) from None return foast.Attribute( - value=new_value, - attr=node.attr, - location=node.location, - type=getattr(new_value.type, node.attr), + value=new_value, attr=node.attr, location=node.location, type=attr_type ) def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscript: @@ -784,20 +793,11 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call: ): raise errors.DSLError(node.location, "Functions can only be called directly.") elif isinstance(new_func.type, ts.FieldType): - for arg in new_args: - # A Cartesian `FieldOffset` shifts by the index it is subscripted with, so it - # carries no displacement on its own. Only an offset with a local dimension is - # meaningful unsubscripted, as the neighbor access `field(Off)`. - if ( - isinstance(arg, (foast.Name, foast.Attribute)) - and isinstance(arg.type, ts.OffsetType) - and len(arg.type.target) == 1 - ): - raise errors.DSLError( - arg.location, - f"Cannot shift by the Cartesian offset '{arg!s}' without an index.", - hints=[f"Give the displacement, e.g. '{arg!s}[1]'."], - ) + # NOTE: a bare single-target offset used to be rejected here, as the unsubscripted + # Cartesian `FieldOffset` `a(Koff)`. There is no such declaration any more: a Cartesian + # shift is `a(Dim + i)`, and the only single-target offset left is the result of + # `as_offset`, which is a call, not a name. + pass elif isinstance(new_func.type, ts.DimensionType): assert new_func.type.dim.kind == DimensionKind.LOCAL return foast.Call( @@ -988,16 +988,14 @@ def _visit_astype(self, node: foast.Call, **kwargs: Any) -> foast.Call: def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: arg_0 = node.args[0].type arg_1 = node.args[1].type - assert isinstance(arg_0, ts.OffsetType) assert isinstance(arg_1, ts.FieldType) - if not fbuiltins.is_cartesian_offset(arg_0): - target_dims = ", ".join(d.__qualname__ for d in arg_0.target) # for the diagnostic + if not isinstance(arg_0, ts.DimensionType) or arg_0.dim.kind is common.DimensionKind.LOCAL: raise errors.DSLError( node.location, - f"'as_offset' is only supported for Cartesian offsets " - f"(single target dimension equal to source dimension); " - f"got source '{arg_0.source.__qualname__}' and target ({target_dims}).", + f"'as_offset' shifts along a non-local dimension, e.g. 'as_offset(KDim, field)';" + f" got '{arg_0}'.", ) + dim = arg_0.dim if not type_info.is_integral(arg_1): raise errors.DSLError( node.location, @@ -1006,16 +1004,20 @@ def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: f"{node.location}", ) - if arg_0.source not in arg_1.dims: + if dim not in arg_1.dims: raise errors.DSLError( node.location, f"Incompatible argument in call to '{node.func!s}': " - f"'{arg_0.source}' not in list of offset field dimensions '{arg_1.dims}'. " + f"'{dim}' not in list of offset field dimensions '{arg_1.dims}'. " f"{node.location}", ) return foast.Call( - func=node.func, args=node.args, kwargs=node.kwargs, type=arg_0, location=node.location + func=node.func, + args=node.args, + kwargs=node.kwargs, + type=ts.OffsetType(source=dim, target=(dim,)), + location=node.location, ) def _deduce_where_return_type( diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 841ce4ca98..d1b0a6cbb1 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -234,12 +234,12 @@ def visit_Symbol(self, node: foast.Symbol, **kwargs: Any) -> itir.Sym: def visit_Name(self, node: foast.Name, **kwargs: Any) -> itir.SymRef | itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag) return im.ref(node.id) def visit_Attribute(self, node: foast.Attribute, **kwargs: Any) -> itir.AxisLiteral: if isinstance(node.type, ts.DimensionType): - return itir.AxisLiteral(value=node.type.dim.tag, kind=node.type.dim.kind) + return itir.AxisLiteral(value=node.type.dim.tag) if isinstance(named_tup_type := node.value.type, ts.NamedCollectionType): ind = named_tup_type.keys.index(node.attr) @@ -331,9 +331,9 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: # `field(as_offset(Off, offset_field))` case foast.Call(func=foast.Name(id="as_offset")): func_args = arg - offset_type = func_args.args[0].type - assert isinstance(offset_type, ts.OffsetType) - dim = offset_type.source + dim_type = func_args.args[0].type + assert isinstance(dim_type, ts.DimensionType) + dim = dim_type.dim offset_field = self.visit(func_args.args[1], **kwargs) current_expr = im.as_fieldop( im.lambda_("__it", "__offset")( diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 1b6cefc6b2..57787f4d1f 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -17,7 +17,6 @@ from gt4py.eve import NodeTranslator, traits from gt4py.next import common, config, errors, utils from gt4py.next.ffront import ( - fbuiltins, gtcallable, program_ast as past, stages as ffront_stages, @@ -74,7 +73,7 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> stages.CompilableProgramDef: """ all_closure_vars = transform_utils._get_closure_vars_recursively(inp.data.closure_vars) offsets_and_dimensions = transform_utils._filter_closure_vars_by_type( - all_closure_vars, fbuiltins.FieldOffset, common.DimensionMeta + all_closure_vars, common.ConnectivityMeta, common.DimensionMeta ) grid_type = transform_utils._deduce_grid_type( inp.data.grid_type, offsets_and_dimensions.values() @@ -383,9 +382,7 @@ def _construct_itir_domain_arg( domain_args = [] for dim_i, dim in enumerate(out_type.dims): # an expression for the range of a dimension - dim_range = im.call("get_domain_range")( - out_expr, itir.AxisLiteral(value=dim.tag, kind=dim.kind) - ) + dim_range = im.call("get_domain_range")(out_expr, itir.AxisLiteral(value=dim.tag)) dim_start, dim_stop = im.tuple_get(0, dim_range), im.tuple_get(1, dim_range) # bounds @@ -413,7 +410,7 @@ def _construct_itir_domain_arg( domain_args.append( itir.FunCall( fun=itir.SymRef(id="named_range"), - args=[itir.AxisLiteral(value=dim.tag, kind=dim.kind), lower, upper], + args=[itir.AxisLiteral(value=dim.tag), lower, upper], ) ) diff --git a/src/gt4py/next/ffront/transform_utils.py b/src/gt4py/next/ffront/transform_utils.py index 09c9d4b9ee..26bfc2f765 100644 --- a/src/gt4py/next/ffront/transform_utils.py +++ b/src/gt4py/next/ffront/transform_utils.py @@ -10,7 +10,6 @@ from typing import Any, Iterable, Optional from gt4py.next import common -from gt4py.next.ffront import fbuiltins from gt4py.next.ffront.gtcallable import GTCallable @@ -47,7 +46,7 @@ def _filter_closure_vars_by_type(closure_vars: dict[str, Any], *types: type) -> def _deduce_grid_type( requested_grid_type: Optional[common.GridType], - offsets_and_dimensions: Iterable[fbuiltins.FieldOffset | common.Dimension], + offsets_and_dimensions: Iterable[type[common.NeighborConnectivity] | common.Dimension], ) -> common.GridType: """ Derive grid type from actually occurring dimensions and check against optional user request. @@ -59,7 +58,7 @@ def _deduce_grid_type( deduced_grid_type = common.GridType.CARTESIAN for o in offsets_and_dimensions: - if isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o): + if isinstance(o, common.ConnectivityMeta): deduced_grid_type = common.GridType.UNSTRUCTURED break if isinstance(o, common.DimensionMeta) and o.kind == common.DimensionKind.LOCAL: @@ -71,7 +70,7 @@ def _deduce_grid_type( and deduced_grid_type == common.GridType.UNSTRUCTURED ): raise ValueError( - "'grid_type == GridType.CARTESIAN' was requested, but unstructured 'FieldOffset' or local 'Dimension' was found." + "'grid_type == GridType.CARTESIAN' was requested, but a 'NeighborConnectivity' or a local dimension was found." ) return deduced_grid_type if requested_grid_type is None else requested_grid_type diff --git a/src/gt4py/next/fingerprinting.py b/src/gt4py/next/fingerprinting.py index b762da1ec9..0ba94993a1 100644 --- a/src/gt4py/next/fingerprinting.py +++ b/src/gt4py/next/fingerprinting.py @@ -231,6 +231,24 @@ def _instance_state(obj: Any) -> tuple[Any, ...]: if "base" in obj.__dict__ else EmptyDeconstruction.from_reference(obj) ), + # A connectivity declaration is a class, fingerprinted by reference like a dimension, *and* + # by what it declares: redefining `V2E` under the same name with other dimensions or counts + # (e.g. re-running a notebook cell) must not reuse artifacts built for the old declaration. + common.ConnectivityMeta: lambda obj: ( + # NOTE: the name goes into the state unverified; the local dimension is fingerprinted as a + # class, so the strict fingerprinter still checks *its* importability (which is the + # connectivity's own, for a nested `Local`, and another module's for a shared one). + Deconstruction.from_pieces( + obj.origin, + obj.codomain, + common.local_dimension_of(obj), + common.local_dimension_of(obj).max_neighbors, + common.local_dimension_of(obj).min_neighbors, + state=b"neighbor_connectivity\0" + obj.tag.encode(), + ) + if "Local" in obj.__dict__ + else EmptyDeconstruction.from_reference(obj) + ), type(None): lambda obj: EmptyDeconstruction.from_typed_value(type(None)), bool: lambda obj: EmptyDeconstruction.from_typed_value(bool, b"1" if obj else b"0"), int: lambda obj: EmptyDeconstruction.from_typed_value(type(obj), str(int(obj)).encode()), diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 4f3ddd732c..6cfebfd529 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -51,7 +51,6 @@ exceptions as embedded_exceptions, operators, ) -from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import builtins, runtime from gt4py.next.type_system import type_specifications as ts, type_translation @@ -154,7 +153,10 @@ def ndarray(self) -> core_defs.NDArrayObject: def asnumpy(self) -> np.ndarray: raise NotImplementedError - def premap(self, index_field: common.Connectivity | fbuiltins.FieldOffset) -> common.Field: + def premap( + self, + index_field: common.Connectivity | type[common.NeighborConnectivity], + ) -> common.Field: raise NotImplementedError def restrict( # type: ignore[override] @@ -171,8 +173,8 @@ def as_scalar(self) -> xtyping.Never: def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError() @@ -214,12 +216,6 @@ def skip_value( NamedFieldIndices: TypeAlias = Mapping[Tag, FieldIndex | SparsePositionEntry] -# Magic local dimension for the result of a `make_const_list`. -# A clean implementation will probably involve to tag the `make_const_list` -# with the neighborhood it is meant to be used with. -_CONST_DIM = common.ConstListDim - - @runtime_checkable class ItIterator(Protocol): """ @@ -568,10 +564,15 @@ def execute_shift( for i, p in reversed(list(enumerate(new_entry))): # first shift applies to the last sparse dimensions of that axis type if p is None: - if tag == _CONST_DIM.tag: + if tag == common.ConstList.tag: new_entry[i] = 0 else: - offset_implementation = common.get_offset(offset_provider, tag) + # NOTE: the sparse tag is the local dimension's; the table over it may be + # keyed by a connectivity sharing it (see `common.connectivity_key_over`). + offset_implementation = common.get_offset( + offset_provider, + common.connectivity_key_over(offset_provider, tag), + ) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().source_dim cur_index = pos[source_dim.tag] @@ -1000,13 +1001,14 @@ def field_getitem(self, named_indices: NamedFieldIndices) -> Any: def field_setitem(self, named_indices: NamedFieldIndices, value: Any): if isinstance(self._ndarrayfield, common.MutableField): if isinstance(value, _List): + local_tag = value.local_dim.tag for i, v in enumerate(value): # type:ignore[var-annotated, arg-type] self._ndarrayfield[ - self._translate_named_indices({**named_indices, value.offset.value: i}) # type: ignore[dict-item] + self._translate_named_indices({**named_indices, local_tag: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ - self._translate_named_indices({**named_indices, _CONST_DIM.tag: 0}) + self._translate_named_indices({**named_indices, common.ConstList.tag: 0}) ] = value.value else: self._ndarrayfield[self._translate_named_indices(named_indices)] = value @@ -1152,8 +1154,8 @@ def as_scalar(self) -> core_defs.IntegralScalar: def premap( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1293,8 +1295,8 @@ def asnumpy(self) -> np.ndarray: def premap( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1389,10 +1391,10 @@ def constant_field(value: Any, dtype_like: Optional[core_defs.DTypeLike] = None) @builtins.shift.register(EMBEDDED) def shift( - *offsets: Union[runtime.Offset, OffsetPart], + *offsets: Union[runtime.Offset, type[common.NeighborConnectivity], OffsetPart], ) -> Callable[[ItIterator], ItIterator]: def impl(it: ItIterator) -> ItIterator: - return it.shift(*list(o.value if isinstance(o, runtime.Offset) else o for o in offsets)) + return it.shift(*list(_as_offset_tag(o) for o in offsets)) return impl @@ -1409,16 +1411,26 @@ def __getitem__(self, i: int): return self.values[i] def __gt_type__(self) -> ts.ListType: - offset_tag = self.offset.value - assert isinstance(offset_tag, str) element_type = type_translation.from_value(self.values[0]) assert isinstance(element_type, ts.DataType) + return ts.ListType(element_type=element_type, offset_type=self.local_dim) + + @property + def local_dim(self) -> common.Dimension: + """ + The local dimension the list runs along. + + The neighbor dimension of the connectivity the list was built with, which is not + necessarily named like its offset: a connectivity can share another one's local + dimension. + """ + offset_tag = self.offset.value + assert isinstance(offset_tag, str) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None connectivity = common.get_offset(offset_provider, offset_tag) assert common.is_neighbor_table(connectivity) - local_dim = connectivity.__gt_type__().neighbor_dim - return ts.ListType(element_type=element_type, offset_type=local_dim) + return connectivity.__gt_type__().neighbor_dim @dataclasses.dataclass(frozen=True) @@ -1433,13 +1445,26 @@ 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, ) +def _as_offset_tag( + offset: runtime.Offset | type[common.NeighborConnectivity] | OffsetPart, +) -> OffsetPart: + if isinstance(offset, common.ConnectivityMeta): + return offset.offset_tag + return offset.value if isinstance(offset, runtime.Offset) else offset + + @builtins.neighbors.register(EMBEDDED) -def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: - offset_str = offset.value if isinstance(offset, runtime.Offset) else offset +def neighbors(offset: runtime.Offset | type[common.NeighborConnectivity], it: ItIterator) -> _List: + field_offset: runtime.Offset = ( + runtime.Offset(value=offset.offset_tag) + if isinstance(offset, common.ConnectivityMeta) + else offset + ) + offset_str = _as_offset_tag(field_offset) assert isinstance(offset_str, str) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None @@ -1451,7 +1476,7 @@ def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: for i in range(connectivity.__gt_type__().max_neighbors) if (shifted := it.shift(offset_str, i)).can_deref() ), - offset=offset, + offset=field_offset, ) @@ -1463,12 +1488,14 @@ def list_get(i, lst: _List[Optional[DT]]) -> Optional[DT] | Undefined: def _get_offset(*lists: _List | _ConstList) -> Optional[runtime.Offset]: - offsets = set((lst.offset for lst in lists if hasattr(lst, "offset"))) - if len(offsets) == 0: + neighbor_lists = [lst for lst in lists if isinstance(lst, _List)] + if len(neighbor_lists) == 0: return None - if len(offsets) == 1: - return offsets.pop() - raise AssertionError("All lists must have the same offset.") + # NOTE: compared by local dimension, not by offset: a connectivity and one sharing its local + # dimension build lists along the same axis. + if len({lst.local_dim for lst in neighbor_lists}) != 1: + raise AssertionError("All lists must run along the same local dimension.") + return neighbor_lists[0].offset @builtins.map_list.register(EMBEDDED) @@ -1515,13 +1542,16 @@ class SparseListIterator: offsets: Sequence[OffsetPart] = dataclasses.field(default_factory=list, kw_only=True) def deref(self) -> Any: - if self.list_offset == _CONST_DIM.tag: + if self.list_offset == common.ConstList.tag: return _ConstList( value=self.it.shift(*self.offsets, SparseTag(self.list_offset), 0).deref() ) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None - connectivity = common.get_offset(offset_provider, self.list_offset) + # NOTE: `list_offset` is the local dimension's tag; the table over it may be keyed by a + # connectivity sharing it (see `common.connectivity_key_over`). + connectivity_key = common.connectivity_key_over(offset_provider, self.list_offset) + connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( values=tuple( @@ -1531,7 +1561,7 @@ def deref(self) -> Any: shifted := self.it.shift(*self.offsets, SparseTag(self.list_offset), i) ).can_deref() ), - offset=runtime.Offset(value=self.list_offset), + offset=runtime.Offset(value=connectivity_key), ) def can_deref(self) -> bool: @@ -1770,15 +1800,17 @@ 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.DimensionMeta) - connectivity = common.get_offset(offset_provider, offset_type.tag) + connectivity = common.get_offset( + offset_provider, common.connectivity_key_over(offset_provider, offset_type) + ) assert common.is_neighbor_table(connectivity) return domain.insert( len(domain), diff --git a/src/gt4py/next/iterator/ir.py b/src/gt4py/next/iterator/ir.py index f024ef6168..05561c2398 100644 --- a/src/gt4py/next/iterator/ir.py +++ b/src/gt4py/next/iterator/ir.py @@ -90,10 +90,19 @@ class OffsetLiteral(Expr): class AxisLiteral(Expr): - # TODO(havogt): Refactor to use declare Axis/Dimension at the Program level. - # Now every use of the literal has to provide the kind, where usually we only care of the name. + #: The dimension's tag, its qualified Python name (ADR 0028). value: str - kind: common.DimensionKind = common.DimensionKind.HORIZONTAL + + @property + def dim(self) -> common.Dimension: + """The dimension the literal names, resolved from its tag.""" + return common.resolve(self.value) + + @property + def kind(self) -> common.DimensionKind: + # NOTE: derived, not stored: the dimension class carries its kind, so a stored copy could + # only disagree with it (it used to, for local dimensions printed as vertical). + return self.dim.kind class CartesianOffset(Expr): diff --git a/src/gt4py/next/iterator/ir_utils/ir_makers.py b/src/gt4py/next/iterator/ir_utils/ir_makers.py index d6458f7a88..053e7b78c8 100644 --- a/src/gt4py/next/iterator/ir_utils/ir_makers.py +++ b/src/gt4py/next/iterator/ir_utils/ir_makers.py @@ -583,7 +583,7 @@ def _impl(*its: itir.Expr) -> itir.FunCall: def axis_literal(dim: common.Dimension) -> itir.AxisLiteral: - return itir.AxisLiteral(value=dim.tag, kind=dim.kind) + return itir.AxisLiteral(value=dim.tag) def broadcast(expr: ExprLike, dims: Iterable[common.Dimension]) -> itir.FunCall: @@ -640,7 +640,7 @@ def index(dim: common.Dimension) -> itir.FunCall: Returns: A function that constructs a Field of indices in the given dimension. """ - return call("index")(itir.AxisLiteral(value=dim.tag, kind=dim.kind)) + return call("index")(itir.AxisLiteral(value=dim.tag)) def map_list(op): diff --git a/src/gt4py/next/iterator/pretty_parser.py b/src/gt4py/next/iterator/pretty_parser.py index 07033ff1db..0324499d26 100644 --- a/src/gt4py/next/iterator/pretty_parser.py +++ b/src/gt4py/next/iterator/pretty_parser.py @@ -20,7 +20,7 @@ from gt4py.next.type_system import type_specifications as ts -GRAMMAR = """ +GRAMMAR = r""" start: fencil_definition | function_definition | declaration @@ -41,7 +41,7 @@ // suffix terminates it. TAG: /[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*(?:\[[A-Za-z_]\w*(?:\.[A-Za-z_]\w*)*\])?/ OFFSET_LITERAL: ( INT_LITERAL | TAG ) "ₒ" - AXIS_LITERAL: TAG ("ᵥ" | "ₕ") + AXIS_LITERAL: TAG ("ᵥ" | "ₕ" | "ₗ") INFINITY_LITERAL: "∞" | "-∞" _literal: INT_LITERAL | FLOAT_LITERAL | OFFSET_LITERAL | AXIS_LITERAL | INFINITY_LITERAL ID_NAME: CNAME @@ -177,9 +177,8 @@ def INFINITY_LITERAL(self, value: lark_lexer.Token) -> ir.InfinityLiteral: return ir.InfinityLiteral.POSITIVE def AXIS_LITERAL(self, value: lark_lexer.Token) -> ir.AxisLiteral: - name = value.value[:-1] - kind = ir.DimensionKind.HORIZONTAL if value.value[-1] == "ₕ" else ir.DimensionKind.VERTICAL - return ir.AxisLiteral(value=name, kind=kind) + # NOTE: the kind suffix is only for the reader; the kind is the dimension's own. + return ir.AxisLiteral(value=value.value[:-1]) def lam(self, *args: ir.Node) -> ir.Lambda: *params, expr = args diff --git a/src/gt4py/next/iterator/pretty_printer.py b/src/gt4py/next/iterator/pretty_printer.py index 5fbba8920e..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/runtime.py b/src/gt4py/next/iterator/runtime.py index 88a466229f..2df378346c 100644 --- a/src/gt4py/next/iterator/runtime.py +++ b/src/gt4py/next/iterator/runtime.py @@ -78,7 +78,11 @@ def __call__( offset_provider=None, column_axis=None, ): - offset_provider = offset_provider or self.offset_provider + # NOTE: not strict: iterator IR names offsets by arbitrary strings. + offset_provider = common.as_tag_keyed_offset_provider( + offset_provider or self.offset_provider, strict=False + ) + common.check_offset_provider(offset_provider) column_axis = column_axis or self.column_axis if backend is not None: diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 4e5b8d4b5a..fb4eeddb37 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -141,7 +141,7 @@ def make_node(o): if isinstance(o, Node): return o if isinstance(o, common.DimensionMeta): - return AxisLiteral(value=o.tag, kind=o.kind) + return AxisLiteral(value=o.tag) if isinstance(o, common.Infinity): if o is common.Infinity.POSITIVE: return itir.InfinityLiteral.POSITIVE @@ -153,6 +153,8 @@ def make_node(o): # it, see `execute_shift`); decide whether to fold it into the shift value or forbid it. assert o.offset == 0 return im.cartesian_offset(o.domain_dim, o.codomain) + if isinstance(o, common.ConnectivityMeta): + return OffsetLiteral(value=o.offset_tag) if callable(o): if o.__name__ == "": return lambdadef(o) diff --git a/src/gt4py/next/iterator/transforms/remove_broadcast.py b/src/gt4py/next/iterator/transforms/remove_broadcast.py index db6445cc51..1a92361216 100644 --- a/src/gt4py/next/iterator/transforms/remove_broadcast.py +++ b/src/gt4py/next/iterator/transforms/remove_broadcast.py @@ -34,9 +34,7 @@ class RemoveBroadcast(PreserveLocationVisitor, NodeTranslator): >>> domain = im.domain(common.GridType.CARTESIAN, {IDim: (0, 10), JDim: (0, 10)}) >>> expr = im.call("broadcast")( ... im.ref("inp"), - ... im.make_tuple( - ... *(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (IDim, JDim)) - ... ), + ... im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (IDim, JDim))), ... ) >>> expr.annex.domain = domain_utils.SymbolicDomain.from_expr(domain) >>> transformed = RemoveBroadcast.apply(expr) diff --git a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py index 3714ea56ad..a9b3a35386 100644 --- a/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py +++ b/src/gt4py/next/iterator/transforms/replace_get_domain_range_with_constants.py @@ -56,6 +56,9 @@ class ReplaceGetDomainRangeWithConstants(PreserveLocationVisitor, NodeTranslator >>> from gt4py import next as gtx >>> class KDim(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... >>> class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... + >>> import sys # register the dimensions where their tags point, as a module would + >>> sys.modules[__name__].KDim = KDim + >>> sys.modules[__name__].Vertex = Vertex >>> sizes = { ... "out": gtx.domain({Vertex: (0, 10), KDim: (0, 20)}), diff --git a/src/gt4py/next/iterator/transforms/unroll_reduce.py b/src/gt4py/next/iterator/transforms/unroll_reduce.py index cfb7bb3226..22491346bb 100644 --- a/src/gt4py/next/iterator/transforms/unroll_reduce.py +++ b/src/gt4py/next/iterator/transforms/unroll_reduce.py @@ -40,11 +40,11 @@ def _get_neighbors_args(reduce_args: Iterable[itir.Expr]) -> Iterator[itir.FunCa return filter(_is_neighbors_or_lifted_and_neighbors, flat_reduce_args) -def _get_partial_offset_tags(reduce_args: Iterable[itir.Expr]) -> Iterable[str]: +def _get_partial_local_dims(reduce_args: Iterable[itir.Expr]) -> Iterable[common.Dimension]: assert all(isinstance(arg.type, ts.ListType) for arg in reduce_args) return [ - arg.type.offset_type.tag # type: ignore[union-attr] # checked in previous lines + arg.type.offset_type # type: ignore[union-attr] # checked in previous lines for arg in reduce_args if arg.type.offset_type is not None # type: ignore[union-attr] # checked in previous lines ] @@ -59,8 +59,10 @@ def _get_connectivity( raise ValueError("Expected a call to a 'reduce' object, i.e. 'reduce(...)(...)'.") connectivities: list[common.NeighborConnectivityType] = [] - for o in _get_partial_offset_tags(applied_reduce_node.args): - conn = common.get_offset_type(offset_provider_type, o) + for local_dim in _get_partial_local_dims(applied_reduce_node.args): + conn = common.get_offset_type( + offset_provider_type, common.connectivity_key_over(offset_provider_type, local_dim) + ) assert isinstance(conn, common.NeighborConnectivityType) connectivities.append(conn) diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 7415bcb7a1..17b0949c6a 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -406,7 +406,7 @@ def _canonicalize_nb_fields( Examples: >>> class Vertex(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> input_field = ts.FieldType( ... dims=[ ... Vertex, diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index 0b49f9ec0c..6852a27619 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -41,7 +41,8 @@ ScalarOrTupleOfScalars: TypeAlias = xtyping.MaybeNestedInTuple[core_defs.Scalar] -#: Content of the key: (*hashable_arg_descriptors, id(offset_provider), concrete_instantation_if_generic) +#: Content of the key: (*hashable_arg_descriptors, hash of the offset provider's (tag, id(table)) +#: items, concrete_instantation_if_generic) CompiledProgramsKey: TypeAlias = tuple[tuple[Hashable, ...], int, str | None] ArgStaticDescriptorsByType: TypeAlias = dict[ @@ -643,6 +644,7 @@ def _compile_variant( else: raise ValueError(f"Invalid 'offset_provider': {offset_provider}") + common.check_offset_provider(offset_provider, deep=True) self._initialize_argument_descriptor_mapping(argument_descriptors) _validate_argument_descriptors(self.program_type, argument_descriptors) diff --git a/src/gt4py/next/otf/options.py b/src/gt4py/next/otf/options.py index 4f77d44586..1056e2da5d 100644 --- a/src/gt4py/next/otf/options.py +++ b/src/gt4py/next/otf/options.py @@ -15,7 +15,7 @@ class CompilationOptionsArgs(TypedDict, total=False): enable_jit: bool static_params: Sequence[str] - connectivities: common.OffsetProvider + connectivities: common.OffsetProviderLike static_domains: bool @@ -36,9 +36,17 @@ class CompilationOptions: #: A dictionary holding static/compile-time information about the offset providers. #: For now, it is used for ahead of time compilation in DaCe orchestrated programs, #: i.e. DaCe programs that call GT4Py Programs -SDFGConvertible interface-. - connectivities: common.OffsetProvider | None = None + connectivities: common.OffsetProviderLike | None = None static_domains: bool = False + def __post_init__(self) -> None: + if self.connectivities is not None: + object.__setattr__( + self, "connectivities", common.as_tag_keyed_offset_provider(self.connectivities) + ) + # the DaCe orchestration reads these directly, without passing an offset provider + common.check_offset_provider(self.connectivities, deep=True) + assert CompilationOptionsArgs.__annotations__.keys() == CompilationOptions.__annotations__.keys() diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index ca67fb18af..c4560a60fe 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -17,7 +17,6 @@ from gt4py._core import definitions as core_defs from gt4py.eve import codegen from gt4py.next import common -from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import ir as itir from gt4py.next.iterator.transforms import pass_manager from gt4py.next.otf import artifacts, stages, workflow @@ -81,17 +80,15 @@ def _process_regular_arguments( if isinstance(parameter.type_, ts.FieldType): for dim in parameter.type_.dims: - if ( - isinstance( - dim, fbuiltins.FieldOffset - ) # TODO(havogt): remove support for FieldOffset as Dimension - or dim.kind is common.DimensionKind.LOCAL - ): + if dim.kind is common.DimensionKind.LOCAL: # translate sparse dimensions to tuple dtype - # NOTE: the tag is the offset-provider key, and its mangled form names the - # `generated::_t` tag type. A legacy `FieldOffset` carries it as `value`. - dim_name = dim.value if isinstance(dim, fbuiltins.FieldOffset) else dim.tag - connectivity = common.get_offset_type(offset_provider_type, dim_name) + # NOTE: the local dimension's tag names the `generated::_t` tag type + # (mangled); its table may be keyed by a connectivity sharing it. + dim_name = dim.tag + connectivity = common.get_offset_type( + offset_provider_type, + common.connectivity_key_over(offset_provider_type, dim), + ) assert isinstance(connectivity, common.NeighborConnectivityType) size = connectivity.max_neighbors arg = f"gridtools::sid::dimension_to_tuple_like({arg})" diff --git a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py index 74684d420b..38fe8fa1c6 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py @@ -254,11 +254,16 @@ def visit_FunCall(self, node: itir.FunCall) -> itir.FunCall: assert isinstance(node.args[0], itir.FunCall) first_axis_literal = node.args[0].args[0] assert isinstance(first_axis_literal, itir.AxisLiteral) - if first_axis_literal.kind == itir.DimensionKind.VERTICAL: + if ir_utils_misc.dim_from_axis_literal(first_axis_literal).kind == ( + itir.DimensionKind.VERTICAL + ): assert len(node.args) == 2 assert isinstance(node.args[1], itir.FunCall) assert isinstance(node.args[1].args[0], itir.AxisLiteral) - assert node.args[1].args[0].kind == itir.DimensionKind.HORIZONTAL + assert ( + ir_utils_misc.dim_from_axis_literal(node.args[1].args[0]).kind + == itir.DimensionKind.HORIZONTAL + ) return itir.FunCall(fun=node.fun, args=[node.args[1], node.args[0]]) return node diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index c4b1526bdc..d3026b3ea2 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -89,6 +89,11 @@ class DataflowBuilder(Protocol): @abc.abstractmethod def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: ... + @abc.abstractmethod + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + """The offset of a connectivity over `local_dim`, see `common.connectivity_key_over`.""" + ... + @abc.abstractmethod def unique_nsdfg_name(self, prefix: str) -> str: ... @@ -564,6 +569,9 @@ class GTIRToSDFG(eve.NodeVisitor, SDFGBuilder): def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: return gtx_common.get_offset_type(self.offset_provider_type, offset) + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + return gtx_common.connectivity_key_over(self.offset_provider_type, local_dim) + def make_field( self, data_node: dace_nodes.AccessNode, @@ -578,10 +586,12 @@ def make_field( # the local dimension is converted into `ListType` data element if not isinstance(data_type.dtype, ts.ScalarType): raise ValueError(f"Invalid field type {data_type}.") - if not gtx_common.has_offset(self.offset_provider_type, local_dim.tag): + try: + self.connectivity_key_over(local_dim) + except KeyError as ex: raise ValueError( f"The provided local dimension {local_dim} does not match any offset provider type." - ) + ) from ex local_type = ts.ListType(element_type=data_type.dtype, offset_type=local_dim) field_type = ts.FieldType( dims=[dim for dim in data_type.dims if dim != local_dim], dtype=local_type @@ -839,7 +849,7 @@ def _make_array_shape_and_strides( for dim in dims: if dim.kind == gtx_common.DimensionKind.LOCAL: # for local dimension, the size is taken from the associated connectivity type - shape.append(neighbor_table_types[dim.tag].max_neighbors) + shape.append(gtx_dace_args.local_dimension_size(name, dim, neighbor_table_types)) elif gtx_dace_args.is_connectivity_identifier(name, self.offset_provider_type): # we use symbolic size for the global dimension of a connectivity shape.append(gtx_dace_args.field_size_symbol(name, dim, neighbor_table_types)) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py index 737648b911..60827eeedc 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py @@ -254,7 +254,9 @@ def translate_concat_where( local_dim = node.type.dtype.offset_type assert local_dim is not None dtype = gtx_dace_args.as_dace_type(node.type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type(local_dim.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(local_dim) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) local_idx = gtx_common.order_dimensions([*output_dims, local_dim]).index(local_dim) output_shape.insert(local_idx, offset_provider_type.max_neighbors) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 924758c727..8a8e3fce53 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -58,13 +58,6 @@ ) -# Magic local dimension used for list of values with length known at compile-time. -# NOTE: the canonical class from `common`, not a local declaration: under nominal identity a -# second declaration would be a *different* dimension and the `== _CONST_DIM` checks below -# would stop matching `ListType`s built by embedded execution. -_CONST_DIM: Final = gtx_common.ConstListDim - - @dataclasses.dataclass(frozen=True) class ValueExpr: """ @@ -595,7 +588,7 @@ def _construct_tasklet_result( return ValueExpr( dc_node=temp_node, gt_dtype=( - ts.ListType(element_type=data_type, offset_type=_CONST_DIM) + ts.ListType(element_type=data_type, offset_type=gtx_common.ConstList) if use_array else data_type ), @@ -769,7 +762,9 @@ def _visit_if_branch_arg( local_dim = arg.gt_dtype.offset_type assert local_dim is not None assert isinstance( - self.subgraph_builder.get_offset_provider_type(local_dim.tag), + self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(local_dim) + ), gtx_common.NeighborConnectivityType, ) # find position of the local dimension in the field layout @@ -1142,7 +1137,8 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: gt_field=ts.FieldType( dims=[conn_type.domain[0]], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1318,10 +1314,12 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: assert isinstance(input_arg.gt_dtype, ts.ListType) assert input_arg.gt_dtype.offset_type is not None offset_type = input_arg.gt_dtype.offset_type - if offset_type == _CONST_DIM: + if offset_type is gtx_common.ConstList: # this input argument is the result of `make_const_list` continue - offset_provider_t = self.subgraph_builder.get_offset_provider_type(offset_type.tag) + offset_provider_t = self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) input_conn_types[offset_type] = offset_provider_t @@ -1359,7 +1357,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: raise ValueError(f"More than one local dimension in map expression {node}.") input_size = input_desc.shape[0] if input_size == 1: - assert input_arg.gt_dtype.offset_type == _CONST_DIM + assert input_arg.gt_dtype.offset_type is gtx_common.ConstList input_memlets[conn] = dace.Memlet(data=input_node.data, subset="0") elif input_size == local_size: input_memlets[conn] = dace.Memlet(data=input_node.data, subset=map_index) @@ -1375,7 +1373,9 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if conn_type.has_skip_values: # In case the `map_list` input expressions contain skip values, we use # the connectivity-based offset provider as mask for map computation. - conn_data = gtx_dace_args.connectivity_identifier(offset_type.tag) + conn_data = gtx_dace_args.connectivity_identifier( + self.subgraph_builder.connectivity_key_over(offset_type) + ) conn_desc = self.sdfg.arrays[conn_data] conn_desc.transient = False @@ -1390,7 +1390,8 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: gt_field=ts.FieldType( dims=[conn_type.domain[0]], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1437,7 +1438,7 @@ def _broadcast_const_list( ) -> ValueExpr: assert list_type.offset_type is not None offset_provider_t = self.subgraph_builder.get_offset_provider_type( - list_type.offset_type.tag + self.subgraph_builder.connectivity_key_over(list_type.offset_type) ) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) local_size = offset_provider_t.max_neighbors @@ -1477,7 +1478,9 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: and input_expr.gt_dtype.offset_type is not None ) offset_type = input_expr.gt_dtype.offset_type - offset_provider_type = self.subgraph_builder.get_offset_provider_type(offset_type.tag) + offset_provider_type = self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) inp_conn = "_in" @@ -1489,7 +1492,9 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: and input_expr.gt_dtype.offset_type is not None ) offset_type = input_expr.gt_dtype.offset_type - connectivity = gtx_dace_args.connectivity_identifier(offset_type.tag) + connectivity = gtx_dace_args.connectivity_identifier( + self.subgraph_builder.connectivity_key_over(offset_type) + ) self.sdfg.arrays[connectivity].transient = False reduce_node = gtx_library_nodes.ReduceWithSkipValues( @@ -1716,7 +1721,8 @@ def _make_unstructured_shift( gt_field=ts.FieldType( dims=[conn_type.source_dim], dtype=ts.ListType( - element_type=tt.from_dtype(conn_type.dtype), offset_type=_CONST_DIM + element_type=tt.from_dtype(conn_type.dtype), + offset_type=gtx_common.ConstList, ), ), subset=dace_subsets.Range.from_string( @@ -1921,7 +1927,7 @@ def _visit_Lambda_impl( and node.expr.type.offset_type is not None and isinstance(result, (MemletExpr, ValueExpr)) and isinstance(result.gt_dtype, ts.ListType) - and result.gt_dtype.offset_type == _CONST_DIM + and result.gt_dtype.offset_type is gtx_common.ConstList ): result = self._broadcast_const_list(result, node.expr.type) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index e265377a45..50a445146c 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -325,7 +325,9 @@ def _construct_if_branch_output( assert out_type.dtype.offset_type is not None assert isinstance(out_type.dtype.element_type, ts.ScalarType) dtype = gtx_dace_args.as_dace_type(out_type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type(out_type.dtype.offset_type.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(out_type.dtype.offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) shape = [*shape, offset_provider_type.max_neighbors] diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index da75325989..fb3099aa3d 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -384,7 +384,9 @@ def get_scan_output_shape( assert isinstance(scan_init_data.gt_type, ts.ListType) assert scan_init_data.gt_type.offset_type offset_type = scan_init_data.gt_type.offset_type - offset_provider_type = sdfg_builder.get_offset_provider_type(offset_type.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) list_size = offset_provider_type.max_neighbors return [scan_column_size, dace.symbolic.SymExpr(list_size)] diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py index 5b0b0601fc..c614719aff 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -111,6 +111,27 @@ def field_stride_symbol( return _field_symbol(field_name, dim, "stride", offset_provider_type) +def local_dimension_size( + field_name: str, + dim: gtx_common.Dimension, + neighbor_table_types: dict[str, gtx_common.NeighborConnectivityType], +) -> int: + """ + Number of neighbors along the local dimension `dim` of the field or connectivity table. + + A connectivity table has its own neighbor count. Any other field finds it in a table over + `dim`: normally the one keyed by `dim`'s tag, but a connectivity sharing `dim` with another + one (see `NeighborConnectivity`) is keyed by its own tag, and may be the only one bound. + """ + if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is not None: + own_type = neighbor_table_types[gtx_common.from_codegen_name(m[1])] + if own_type.neighbor_dim == dim: + return own_type.max_neighbors + return neighbor_table_types[ + gtx_common.connectivity_key_over(neighbor_table_types, dim) + ].max_neighbors + + def _range_symbol_name(field_name: str, dim: gtx_common.Dimension) -> str: """Common part of the name for the range start/stop symbols.""" field_range = im.call("get_domain_range")(field_name, dim) diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py index 1af014470c..d21c475723 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py @@ -92,14 +92,17 @@ def _get_args(sdfg: dace.SDFG, args: Sequence[Any]) -> dict[str, Any]: def get_sdfg_conn_args( sdfg: dace.SDFG, - offset_provider: gtx_common.OffsetProvider, + offset_provider: gtx_common.OffsetProviderLike, ) -> dict[str, core_defs.NDArrayObject]: """ Extracts the connectivity tables that are used in the sdfg and ensures that the memory buffers are allocated for the target device. """ connectivity_args = {} - for offset, connectivity in offset_provider.items(): + # NOTE: not strict, like the other IR-level hooks: the keys of a hand-written program are its + # own business, and a declaration is normalized to its tag either way. + provider = gtx_common.as_tag_keyed_offset_provider(offset_provider, strict=False) + for offset, connectivity in provider.items(): name = gtx_dace_args.connectivity_identifier(offset) if name in sdfg.arrays: assert gtx_common.is_neighbor_table(connectivity) diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 2ce4672a77..da0cf84b1a 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -408,7 +408,7 @@ def is_local_field(type_: ts.FieldType) -> bool: Examples: >>> class V(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> is_local_field( ... ts.FieldType(dims=[V, V2E], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64)) ... ) @@ -586,7 +586,7 @@ def promote( >>> promoted.dims == [I, J, K] and promoted.dtype == dtype True - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> list_dtype = ts.ListType(element_type=dtype, offset_type=V2E) >>> promote( ... ts.FieldType(dims=[I], dtype=list_dtype), diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 50f451fd39..afe5b350da 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -71,7 +71,9 @@ def __str__(self) -> str: class OffsetType(TypeSpec): - # TODO(havogt): replace by ConnectivityType + # NOTE: kept, against the TODO that stood here: since ADR 0029 this types a connectivity + # declaration (`V2E.__gt_type__()`) and the result of `as_offset`, and a `ConnectivityType` + # is what the *bound table* produces. Renaming it would be churn with no user-visible gain. source: common.Dimension target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] #: The offset-provider key; `None` for the untagged Cartesian `Dim + offset`. @@ -117,6 +119,10 @@ class ListType(DataType): """ element_type: DataType + #: The local dimension the list runs along. `None` where type inference does not know it, + #: which is how it spells the result of `make_const_list`; embedded execution and the DaCe + #: lowering use `common.ConstList` for the same thing. + #: TODO(egparedes): use `common.ConstList` in type inference too, and drop `None`. offset_type: common.Dimension | None diff --git a/tests/next_tests/definitions.py b/tests/next_tests/definitions.py index 86a4229860..eb659525ed 100644 --- a/tests/next_tests/definitions.py +++ b/tests/next_tests/definitions.py @@ -98,10 +98,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): USES_INDEX_FIELDS = "uses_index_fields" USES_LIFT = "uses_lift" USES_NEGATIVE_MODULO = "uses_negative_modulo" -USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM = "uses_offset_tag_differing_from_local_dim" -USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION = ( - "uses_offset_tag_differing_from_local_dim_in_reduction" -) USES_ORIGIN = "uses_origin" USES_REDUCE_WITH_LAMBDA = "uses_reduce_with_lambda" USES_SCAN = "uses_scan" @@ -142,10 +138,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): #: the connectivity is looked up in the offset provider by the *local dimension's* name. #: Lifted for the gtfn shift path by #1789; see #: `regression_tests/ffront_tests/test_offset_dimensions_names.py`. -OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE = ( - "'{marker}': '{backend}' looks the connectivity up by the local dimension's name," - " so it must equal the offset tag" -) # Index-only vs. consequential markers: # A `uses_*` marker only affects execution if it appears in one of the skip lists below (and thus # in `BACKEND_SKIP_TEST_MATRIX`); such a marker is "consequential" -- it applies the listed @@ -181,16 +173,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_SCAN_IN_STENCIL, XFAIL, BINDINGS_UNSUPPORTED_MESSAGE), (USES_SPARSE_FIELDS, XFAIL, UNSUPPORTED_MESSAGE), (USES_TUPLE_ITERATOR, XFAIL, UNSUPPORTED_MESSAGE), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] ) EMBEDDED_SKIP_LIST = [ @@ -202,11 +184,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): ), # we can't extract the field type from scan args (EMBEDDED_CONCAT_WHERE_INFINITE_DOMAIN, XFAIL, UNSUPPORTED_MESSAGE), (EMBEDDED_CONCAT_WHERE_NON_CONTIGUOUS_DOMAIN, XFAIL, UNSUPPORTED_MESSAGE), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] JAX_EMBEDDED_SKIP_LIST = EMBEDDED_SKIP_LIST + [ (USES_PROGRAM_WITH_SLICED_OUT_ARGUMENTS, XFAIL, UNSUPPORTED_MESSAGE), @@ -217,15 +194,7 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_TUPLES_ARGS_WITH_DIFFERENT_BUT_PROMOTABLE_DIMS, XFAIL, UNSUPPORTED_MESSAGE), (USES_CONCAT_WHERE, XFAIL, UNSUPPORTED_MESSAGE), ] -GTIR_EMBEDDED_SKIP_LIST = ROUNDTRIP_SKIP_LIST + [ - # NOTE: not in `ROUNDTRIP_SKIP_LIST`: the roundtrip backend passes this, only the - # lower-level `iterator/embedded.py` execution keys on the local dimension's name. - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), -] +GTIR_EMBEDDED_SKIP_LIST = ROUNDTRIP_SKIP_LIST GTFN_SKIP_TEST_LIST = ( COMMON_SKIP_TEST_LIST + DOMAIN_INFERENCE_SKIP_LIST @@ -236,12 +205,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_STRIDED_NEIGHBOR_OFFSET, XFAIL, BINDINGS_UNSUPPORTED_MESSAGE), # max_over broken, see https://github.com/GridTools/gt4py/issues/1289 (USES_MAX_OVER, XFAIL, UNSUPPORTED_MESSAGE), - # NOTE: only the reduction; #1789 lifted this for the shift path. - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] ) diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 464bf7fa1f..3f63cec788 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -35,7 +35,17 @@ # Both modules used to declare their own `Dimension("Vertex")` etc., which compared equal; under # nominal identity (ADR 0028) that would be two different dimensions, and tests that mix a # `toy_connectivity` connectivity with a `cases_utils` mesh would silently stop matching. -from next_tests.toy_connectivity import C2EDim, Cell, E2VDim, Edge, V2EDim, Vertex +from next_tests.toy_connectivity import ( + C2E, + C2EDim, + Cell, + E2V, + E2VDim, + Edge, + V2E, + V2EDim, + Vertex, +) __all__ = [ @@ -52,7 +62,6 @@ "Vertex", "Edge", "Cell", - "EdgeOffset", "MeshDescriptor", "CartesianGridDescriptor", ] @@ -177,20 +186,12 @@ class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... KHalfDim = common.flip_staggered(KDim) -Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) -Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) - - -EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) - -class C2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2V(gtx.NeighborConnectivity[Cell, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) -C2V = gtx.FieldOffset(C2VDim.tag, source=Vertex, target=(Cell, C2VDim)) +C2VDim = C2V.Local size = 10 @@ -222,7 +223,10 @@ def simple_cartesian_grid( name="simple_cartesian_grid", sizes=sizes, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + offset_provider_type=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -308,28 +312,28 @@ def simple_mesh(allocator) -> MeshDescriptor: e2v_arr = np.asarray(e2v_arr, dtype=gtx.IndexType) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 4}, codomain=Edge, data=v2e_arr, skip_value=None, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 4}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 4}, codomain=Edge, data=c2e_arr, @@ -344,7 +348,10 @@ def simple_mesh(allocator) -> MeshDescriptor: num_edges=np.int32(num_edges), num_cells=num_cells, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + offset_provider_type=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -403,28 +410,28 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 5}, codomain=Edge, data=v2e_arr, skip_value=common._DEFAULT_SKIP_VALUE, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 3}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 3}, codomain=Edge, data=c2e_arr, @@ -439,7 +446,10 @@ def skip_value_mesh(allocator) -> MeshDescriptor: num_edges=num_edges, num_cells=num_cells, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + offset_provider_type=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -452,3 +462,18 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) def mesh_descriptor(request, exec_alloc_descriptor) -> MeshDescriptor: yield request.param(exec_alloc_descriptor.allocator) + + +def ir_level(mesh: MeshDescriptor) -> MeshDescriptor: + """ + A copy of `mesh` whose offset provider is keyed by tags, as the IR-level APIs expect. + + User-facing entry points normalize a class-keyed provider themselves; tests that drive the + lowering or the backends directly have to hand them the tag-keyed form. + """ + return types.SimpleNamespace( + **{ + **vars(mesh), + "offset_provider": common.as_tag_keyed_offset_provider(mesh.offset_provider), + } + ) diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py index fa65a38b40..f03f970507 100644 --- a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py +++ b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py @@ -42,10 +42,11 @@ class Cell(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2EDim = C2E.Local C2E_TABLE = np.array( [ diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py index 1e65cd734c..c156a61407 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py @@ -19,8 +19,6 @@ cartesian_case, ) from next_tests.integration_tests.cases_utils import ( - Ioff, - Koff, exec_alloc_descriptor, ) @@ -62,10 +60,10 @@ def test_offset_field(cartesian_case): @gtx.field_operator def testee(a: cases.IKField, offset_field: cases.IKField) -> gtx.Field[[IDim, KDim], bool]: - a_i = a(as_offset(Ioff, offset_field)) + a_i = a(as_offset(IDim, offset_field)) # note: this leads to an access to offset_field in # IDim: (0, out.size[I]), KDim: (0, out.size[K]+1) - a_i_k = a_i(as_offset(Koff, offset_field)) + a_i_k = a_i(as_offset(KDim, offset_field)) b_i = a(IDim + 1) b_i_k = b_i(KDim + 1) return a_i_k == b_i_k @@ -97,7 +95,7 @@ def test_offset_field_of_chained_ops(cartesian_case): def testee(a: cases.IKField, offset_field: cases.IKField) -> cases.IKField: b = a + 1 c = b * 2 - return c(as_offset(Koff, offset_field)) + return c(as_offset(KDim, offset_field)) out = cases.allocate(cartesian_case, testee, cases.RETURN)() a = cases.allocate(cartesian_case, testee, "a").extend({KDim: (0, 1)})() diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py index 4a642f5cc4..68c2dcf72b 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py @@ -252,7 +252,7 @@ def test_compile_unstructured(unstructured_case, compile_testee_unstructured): compile_testee_unstructured(*args, offset_provider=unstructured_case.offset_provider, **kwargs) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), @@ -317,7 +317,7 @@ def test_compile_unstructured_for_two_offset_providers( *args, offset_provider=unstructured_case.offset_provider, **kwargs ) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py index dfe013e45a..365de7d473 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py @@ -454,7 +454,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: t = concat_where(Vertex < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -477,7 +477,7 @@ def testee( t = concat_where(KDim < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() k_mask = np.arange(unstructured_case_3d.default_sizes[KDim]) < 2 cases.verify_with_default_data( unstructured_case_3d, @@ -497,7 +497,7 @@ def test_with_local_and_nonlocal_field(unstructured_case, static_domains: bool): def testee(a: cases.EField, b: cases.VField) -> cases.VField: return neighbor_sum(concat_where(Vertex < 2, a(V2E), b), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -524,7 +524,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), b(V2E)), (c(V2E), d(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -557,7 +557,7 @@ def testee(a: cases.EField) -> tuple[cases.VField, cases.VField]: neighbor_sum(concat_where(Vertex < 2, 3, a(V2E)), axis=V2EDim), ) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -590,7 +590,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), c), (3, b(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py index 753015f0d7..b65a16ffea 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py @@ -31,11 +31,11 @@ def testee( ) # multiplication with shifted `ones` because reduction of only non-shifted field with local dimension is not supported inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) ones = cases.allocate(unstructured_case, testee, "ones").strategy(cases.ConstInitializer(1))() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify( unstructured_case, testee, @@ -55,7 +55,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return inp[V2EDim(0)] + inp[V2EDim(1)] + inp[V2EDim(2)] + inp[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -77,7 +77,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int64 return inp_64[V2EDim(0)] + inp_64[V2EDim(1)] + inp_64[V2EDim(2)] + inp_64[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -99,7 +99,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return neighbor_sum(inp, axis=V2EDim) inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -107,7 +107,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 testee, inp, out=cases.allocate(unstructured_case, testee, cases.RETURN)(), - ref=np.sum(unstructured_case.offset_provider[V2EDim.tag].asnumpy(), axis=1), + ref=np.sum(unstructured_case.offset_provider[V2E].asnumpy(), axis=1), ) @@ -119,7 +119,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: return inp(V2E) out = unstructured_case.as_field( - [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider[V2EDim.tag].asnumpy()) + [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider[V2E].asnumpy()) ) inp = cases.allocate(unstructured_case, testee, "inp")() cases.verify( @@ -127,5 +127,5 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: testee, inp, out=out, - ref=inp.asnumpy()[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], + ref=inp.asnumpy()[unstructured_case.offset_provider[V2E].asnumpy()], ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py index 262751f234..dbc623f2de 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py @@ -60,7 +60,7 @@ def test_import_offset_module_unstructured_shift(unstructured_case): def testee(a: cases.EField) -> cases.VField: return neighbor_sum(a(cases.V2E), axis=cases.V2EDim) - v2e_table = unstructured_case.offset_provider[cases.V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[cases.V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -77,7 +77,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider[cases.E2VDim.tag].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[cases.E2V].asnumpy()[:, 0]], ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py index 39aee5de06..6ceead30d6 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py @@ -483,7 +483,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) ) @@ -641,7 +641,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py new file mode 100644 index 0000000000..07793c44fe --- /dev/null +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py @@ -0,0 +1,195 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +"""A `NeighborConnectivity` declaration used directly in DSL code, on every backend.""" + +import dataclasses +import typing + +import numpy as np +import pytest + +import gt4py.next as gtx +from gt4py.next import Dims, Field, common, constructors, neighbor_sum + +from next_tests import definitions as test_defs +from next_tests.integration_tests import cases, cases_utils +from next_tests.integration_tests.cases_utils import ( # noqa: F401 [unused-import] # fixture + exec_alloc_descriptor, +) + + +class V(gtx.DimensionIndex): ... + + +class E(gtx.DimensionIndex): ... + + +class V2E(gtx.NeighborConnectivity[V, E], max_neighbors=4, min_neighbors=4): + class Local(gtx.LocalDimensionIndex): ... + + +#: A second connectivity over the same neighbor axis, bound to a different table. +class V2EShared(gtx.NeighborConnectivity[V, E]): + Local: typing.TypeAlias = V2E.Local + + +@pytest.fixture +def case(exec_alloc_descriptor): + mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) + v2e_arr = mesh.offset_provider[cases_utils.V2E].asnumpy() + table = constructors.as_connectivity( + domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, + codomain=E, + data=v2e_arr, + skip_value=None, + allocator=exec_alloc_descriptor.allocator, + ) + common.check_neighbor_table(V2E, table) + shared_table = constructors.as_connectivity( + domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, + codomain=E, + data=np.ascontiguousarray(v2e_arr[:, ::-1]), + skip_value=None, + allocator=exec_alloc_descriptor.allocator, + ) + common.check_neighbor_table(V2EShared, shared_table) + return cases.Case( + ( + None + if isinstance(exec_alloc_descriptor, test_defs.EmbeddedDummyBackend) + else exec_alloc_descriptor + ), + offset_provider={V2E: table, V2EShared: shared_table}, + default_sizes={V: mesh.num_vertices, E: mesh.num_edges, V2E.Local: v2e_arr.shape[1]}, + grid_type=common.GridType.UNSTRUCTURED, + allocator=exec_alloc_descriptor.allocator, + ) + + +def _table(case: cases.Case, connectivity=V2E) -> np.ndarray: + return case.offset_provider[connectivity].asnumpy() + + +@pytest.mark.uses_unstructured_shift +def test_shift(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2E[1]) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case)[:, 1]]) + + +@pytest.mark.uses_unstructured_shift +def test_reduction(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return neighbor_sum(a(V2E), axis=V2E.Local) + + cases.verify_with_default_data(case, testee, lambda a: np.sum(a[_table(case)], axis=1)) + + +@pytest.mark.uses_unstructured_shift +def test_sparse_argument(case): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + return neighbor_sum(s * a(V2E), axis=V2E.Local) + + cases.verify_with_default_data(case, testee, lambda s, a: np.sum(s * a[_table(case)], axis=1)) + + +@pytest.mark.uses_unstructured_shift +def test_program(case): + @gtx.field_operator + def shift_by_one(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2E[0]) + + @gtx.program + def testee(a: Field[Dims[E], float], out: Field[Dims[V], float]): + shift_by_one(a, out=out) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case)[:, 0]]) + + +@pytest.mark.uses_unstructured_shift +def test_shift_through_a_shared_local_dimension(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2EShared[1]) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case, V2EShared)[:, 1]]) + + +@pytest.mark.uses_unstructured_shift +def test_reduction_through_a_shared_local_dimension(case): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + # combines a sparse field on the shared axis with each connectivity's neighbors + return neighbor_sum(s * a(V2EShared) - a(V2E), axis=V2E.Local) + + cases.verify_with_default_data( + case, + testee, + lambda s, a: np.sum(s * a[_table(case, V2EShared)] - a[_table(case)], axis=1), + ) + + +@pytest.fixture +def case_without_owner(case): + """Only the sharing connectivity is bound: enough for a shift, which needs only its table.""" + return dataclasses.replace(case, offset_provider={V2EShared: case.offset_provider[V2EShared]}) + + +@pytest.mark.uses_unstructured_shift +def test_shift_through_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2EShared[1]) + + cases.verify_with_default_data( + case_without_owner, testee, lambda a: a[_table(case_without_owner, V2EShared)[:, 1]] + ) + + +@pytest.mark.uses_unstructured_shift +def test_reduction_through_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + return neighbor_sum(s * a(V2EShared), axis=V2E.Local) + + cases.verify_with_default_data( + case_without_owner, + testee, + lambda s, a: np.sum(s * a[_table(case_without_owner, V2EShared)], axis=1), + ) + + +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_if_stmts +def test_if_over_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee(a: Field[Dims[E], float], flag: bool) -> Field[Dims[V], float]: + if flag: + s = a(V2EShared) + else: + s = a(V2EShared) * 2.0 + return neighbor_sum(s, axis=V2E.Local) + + cases.verify_with_default_data( + case_without_owner, + testee, + lambda a, flag: ( + np.sum(a[_table(case_without_owner, V2EShared)], axis=1) * (1.0 if flag else 2.0) + ), + ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py index d882a88ec4..abad348c92 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py @@ -60,7 +60,7 @@ def shift_by_one(in_field: cases.IFloatField) -> cases.IFloatField: # direct call to field operator # TODO(tehrengruber): slicing located fields not supported currently - # shift_by_one(in_field, out=out_field[:-1], offset_provider={"Ioff": IDim}) + # shift_by_one(in_field, out=out_field[:-1]) @gtx.program def shift_by_one_program(in_field: cases.IFloatField, out_field: cases.IFloatField): diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py index dd0cc6fb43..f0ce221851 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py @@ -53,7 +53,7 @@ def testee(edge_f: cases.EField) -> cases.VField: inp = cases.allocate(unstructured_case, testee, "edge_f", strategy=strategy)() out = cases.allocate(unstructured_case, testee, cases.RETURN)() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() ref = np.max( inp.asnumpy()[v2e_table], axis=1, @@ -70,7 +70,7 @@ def minover(edge_f: cases.EField) -> cases.VField: out = min_over(edge_f(V2E), axis=V2EDim) return out - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, minover, @@ -100,7 +100,7 @@ def reduction_ek_field( "fop", [reduction_e_field, reduction_ek_field], ids=lambda fop: fop.__name__ ) def test_neighbor_sum(unstructured_case_3d, fop): - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() edge_f = cases.allocate(unstructured_case_3d, fop, "edge_f")() @@ -152,7 +152,7 @@ def fencil_op(edge_f: EKField) -> VKField: def fencil(edge_f: EKField, out: VKField): fencil_op(edge_f, out=out) - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() field = cases.allocate(unstructured_case_3d, fencil, "edge_f", sizes={KDim: 2})() out = cases.allocate(unstructured_case_3d, fencil_op, cases.RETURN, sizes={KDim: 1})() @@ -185,7 +185,7 @@ def reduce_expr(edge_f: cases.EField) -> cases.VField: def fencil(edge_f: cases.EField, out: cases.VField): reduce_expr(edge_f, out=out) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, fencil, @@ -207,7 +207,7 @@ def test_reduction_with_common_expression(unstructured_case): def testee(flux: cases.EField) -> cases.VField: return neighbor_sum(flux(V2E) + flux(V2E), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -223,7 +223,7 @@ def test_reduction_expression_with_where(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, inp(V2E), inp(V2E)), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -252,7 +252,7 @@ def test_reduction_expression_with_where_and_tuples(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, (inp(V2E), inp(V2E)), (inp(V2E), inp(V2E)))[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -281,7 +281,7 @@ def test_reduction_expression_with_where_and_scalar(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(inp(V2E) + where(mask, inp(V2E), 1), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -326,7 +326,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]], ) @@ -347,7 +347,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() neighbor_0_iter = iter(enumerate(e2v_table[:, 0])) edge_start = next(i for i, v in neighbor_0_iter if v >= ORIGIN) edge_stop = next(i for i, v in neighbor_0_iter if v < ORIGIN) @@ -392,16 +392,16 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_flat, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], ) cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_intermediate_result, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], comparison=lambda inp, tmp: np.all(inp == tmp), ) @@ -409,8 +409,8 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], ) @@ -432,7 +432,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() neighbor_iter = iter(enumerate(e2v_table)) edge_start = next(i for i, v in neighbor_iter if all(v >= ORIGIN)) edge_stop = next(i for i, v in neighbor_iter if any(v < ORIGIN)) @@ -453,12 +453,11 @@ def testee(a: cases.VField) -> cases.VField: unstructured_case, testee, ref=lambda a: np.sum( - np.sum(a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()], axis=1, initial=0)[ - unstructured_case.offset_provider[V2EDim.tag].asnumpy() + np.sum(a[unstructured_case.offset_provider[E2V].asnumpy()], axis=1, initial=0)[ + unstructured_case.offset_provider[V2E].asnumpy() ], axis=1, - where=unstructured_case.offset_provider[V2EDim.tag].asnumpy() - != common._DEFAULT_SKIP_VALUE, + where=unstructured_case.offset_provider[V2E].asnumpy() != common._DEFAULT_SKIP_VALUE, ), comparison=lambda a, tmp_2: np.all(a == tmp_2), ) @@ -479,8 +478,8 @@ def testee(inp: cases.EField) -> cases.EField: unstructured_case, testee, ref=lambda inp: np.sum( - np.sum(inp[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1)[ - unstructured_case.offset_provider[E2VDim.tag].asnumpy() + np.sum(inp[unstructured_case.offset_provider[V2E].asnumpy()], axis=1)[ + unstructured_case.offset_provider[E2V].asnumpy() ], axis=1, ), @@ -497,7 +496,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: tmp = red(E2V[0]) return tmp - v2e = unstructured_case.offset_provider[V2EDim.tag] + v2e = unstructured_case.offset_provider[V2E] cases.verify_with_default_data( unstructured_case, reduce_tuple_element, @@ -506,7 +505,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: axis=1, initial=0, where=v2e.asnumpy() != common._DEFAULT_SKIP_VALUE, - )[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], + )[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]], ) @@ -518,7 +517,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: tmp = neighbor_sum(b(V2E) if 2 < 3 else a(V2E), axis=V2EDim) return tmp - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -540,7 +539,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex], int32]: inp = cases.allocate(unstructured_case, testee, "inp")() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify( unstructured_case, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py index ac348266f3..15a94f2e36 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py @@ -26,7 +26,7 @@ unstructured_case, unstructured_case_3d, ) -from next_tests.integration_tests.cases_utils import Koff, exec_alloc_descriptor, mesh_descriptor +from next_tests.integration_tests.cases_utils import exec_alloc_descriptor, mesh_descriptor @pytest.mark.uses_cartesian_shift @@ -156,7 +156,7 @@ def test_cartesian_half_shift_as_offset(cartesian_case): def testee( a: gtx.Field[[IDim, KHalfDim], np.int32], offset_field: cases.IKField ) -> cases.IKField: - return a(KDim - 0.5)(as_offset(Koff, offset_field)) + return a(KDim - 0.5)(as_offset(KDim, offset_field)) ksize = cartesian_case.default_sizes[KDim] a = cases.allocate(cartesian_case, testee, "a", sizes={KHalfDim: ksize + 1})() @@ -180,7 +180,7 @@ def testee( a: gtx.Field[[IDim, KHalfDim], np.int32], offset_field: cases.IKField ) -> cases.IKField: b = a + 1 - return b(KDim - 0.5)(as_offset(Koff, offset_field)) + return b(KDim - 0.5)(as_offset(KDim, offset_field)) ksize = cartesian_case.default_sizes[KDim] a = cases.allocate(cartesian_case, testee, "a", sizes={KHalfDim: ksize + 1})() @@ -208,7 +208,7 @@ def testee( a: gtx.Field[[Vertex, KHalfDim], np.int32], offset_field: gtx.Field[[Edge, KDim], np.int32], ) -> gtx.Field[[Edge, KDim], np.int32]: - return a(E2V[0])(KDim - 0.5)(as_offset(Koff, offset_field)) + return a(E2V[0])(KDim - 0.5)(as_offset(KDim, offset_field)) nvertices = unstructured_case_3d.default_sizes[Vertex] ksize = unstructured_case_3d.default_sizes[KDim] @@ -223,7 +223,7 @@ def testee( )() out = cases.allocate(unstructured_case_3d, testee, cases.RETURN)() - e2v_table = unstructured_case_3d.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case_3d.offset_provider[E2V].asnumpy() cases.verify( unstructured_case_3d, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py index 0eb876abde..20ef5e789a 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py @@ -83,9 +83,7 @@ def test_verification(testee, exec_alloc_descriptor, mesh_descriptor): a = cases.allocate(unstructured_case, testee, "a")() out = cases.allocate(unstructured_case, testee, "out")() - first_nbs, second_nbs = ( - mesh_descriptor.offset_provider[E2VDim.tag].asnumpy()[:, i] for i in [0, 1] - ) + first_nbs, second_nbs = (mesh_descriptor.offset_provider[E2V].asnumpy()[:, i] for i in [0, 1]) ref = (a.ndarray * 2)[first_nbs] + (a.ndarray * 2)[second_nbs] cases.verify( @@ -105,7 +103,7 @@ def test_temporary_symbols(testee, mesh_descriptor): gtir_with_tmp = apply_common_transforms( testee.gtir, extract_temporaries=True, - offset_provider=mesh_descriptor.offset_provider, + offset_provider=common.as_tag_keyed_offset_provider(mesh_descriptor.offset_provider), ) params = ["num_vertices", "num_edges", "num_cells"] diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index c55a145314..ed91a34e9e 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -164,8 +164,8 @@ def testee(a: cases.EField, b: cases.EField) -> tuple[cases.VField, cases.VField unstructured_case, testee, ref=lambda a, b: [ - np.sum(a[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), - np.sum(b[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), + np.sum(a[unstructured_case.offset_provider[V2E].asnumpy()], axis=1), + np.sum(b[unstructured_case.offset_provider[V2E].asnumpy()], axis=1), ], comparison=lambda a, tmp: (np.all(a[0] == tmp[0]), np.all(a[1] == tmp[1])), ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py index 6a160cefc3..2e4f3d8b78 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py @@ -51,7 +51,7 @@ def testee(a: gtx.Field[[Vertex], np.float64]) -> gtx.Field[[Edge], int64]: tmp = astype(a(E2V), int64) return neighbor_sum(tmp, axis=E2VDim) - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py index b16461447a..094686be80 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py @@ -59,7 +59,7 @@ class Node(gtx.DimensionIndex): ... -class NeighDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class NeighDim(gtx.LocalDimensionIndex): ... def array_maker(*lists): diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py index bd675b5f51..626e83d075 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py @@ -17,7 +17,7 @@ from gt4py.next.iterator.embedded import StridedConnectivityField -class Dummy(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class Dummy(gtx.LocalDimensionIndex): ... class LocA(gtx.DimensionIndex): ... diff --git a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py index 01e4ed92bb..eabf67327d 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py @@ -546,7 +546,7 @@ def test_program_unstructured(unstructured_case): unstructured_case.default_sizes[Cell], unstructured_case.default_sizes[Edge], inout=(out_a_shifted, out_a), - ref=((a.ndarray)[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]], a), + ref=((a.ndarray)[unstructured_case.offset_provider[C2E].asnumpy()[:, 1]], a), ) @@ -600,7 +600,7 @@ def test_program_temporary(unstructured_case): extend={Cell: (-restrict_cell[0], restrict_cell[1])}, )() - e2v = (a.ndarray)[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 1]] + e2v = (a.ndarray)[unstructured_case.offset_provider[E2V].asnumpy()[:, 1]] cases.verify( unstructured_case, prog_temporary, @@ -616,7 +616,7 @@ def test_program_temporary(unstructured_case): inout=(out_edge, out_cell), ref=( e2v[restrict_edge[0] : edge_size + restrict_edge[1]], - e2v[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]][ + e2v[unstructured_case.offset_provider[C2E].asnumpy()[:, 1]][ restrict_cell[0] : cell_size + restrict_cell[1] ], ), diff --git a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py index 4fbb01c72d..7051170dd7 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py @@ -39,11 +39,7 @@ # NOTE: imported, not redeclared. Under nominal identity (ADR 0028) a same-named declaration # here would be a different dimension from the one `toy_connectivity` declares, where the old # `Dimension("...")` values compared equal -- and tests mix objects from both modules. -from next_tests.toy_connectivity import E2VDim, Edge, V2EDim, Vertex - - -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) +from next_tests.toy_connectivity import E2V, E2VDim, Edge, V2E, V2EDim, Vertex def assert_close(expected, actual): diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py index fdf8cf5114..d055f1f094 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py @@ -434,3 +434,52 @@ def test_sparse_shifted_stencil_reduce(program_processor): if validate: assert np.allclose(out.asnumpy(), ref) + + +class V2EShared(gtx.NeighborConnectivity[Vertex, Edge]): + """Shares `V2E`'s local dimension; bound to the table with its columns reversed.""" + + Local = V2EDim + + +v2e_shared_arr = np.ascontiguousarray(v2e_arr[:, ::-1]) +v2e_shared_conn = gtx.as_connectivity( + domain={Vertex: v2e_shared_arr.shape[0], V2EDim: v2e_shared_arr.shape[1]}, + codomain=Edge, + data=v2e_shared_arr, +) + + +@fundef +def shift_through_sharer(in_edges): + return deref(shift(V2EShared, 1)(in_edges)) + + +@fundef +def owner_times_sharer(in_edges): + return reduce(plus, 0)( + map_list(multiplies)(neighbors(V2EShared, in_edges), neighbors(V2E, in_edges)) + ) + + +@pytest.mark.parametrize( + "stencil, ref", + [ + (shift_through_sharer, v2e_shared_arr[:, 1]), + (owner_times_sharer, np.sum(v2e_shared_arr * v2e_arr, axis=1)), + ], +) +def test_connectivity_sharing_a_local_dimension(program_processor, stencil, ref): + program_processor, validate = program_processor + inp = edge_index_field() + out = gtx.as_field([Vertex], np.zeros([9], dtype=inp.dtype)) + + run_processor( + stencil[{Vertex: range(0, 9)}], + program_processor, + inp, + out=out, + offset_provider={V2E.offset_tag: v2e_conn, V2EShared.offset_tag: v2e_shared_conn}, + ) + if validate: + assert np.allclose(out.asnumpy(), ref) diff --git a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py index 689d0d5f71..a33f4d0f12 100644 --- a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py +++ b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py @@ -7,19 +7,13 @@ # SPDX-License-Identifier: BSD-3-Clause """ -Regression tests for the four independently authored names of one connectivity. +Regression tests for the names under which one connectivity is used. -Using a single connectivity requires four strings to agree, none of which is -checked against the others at declaration time: - - N1 the `FieldOffset` tag `FieldOffset("V2E", ...)` - N2 the Python variable it is bound to `V2E = FieldOffset(...)` - N3 the local dimension's name `Dimension("V2E", kind=LOCAL)` - N4 the offset-provider key `offset_provider={"V2E": ...}` - -The `V2EDim = Dimension("V2E")` convention makes all four equal, which hides -which one each execution path actually uses. These tests break the convention -deliberately, one name at a time, so the real requirement is visible. +With `FieldOffset`, using a connectivity required four independently authored strings to +agree -- the offset tag, the Python variable it was bound to, the local dimension's name and +the offset-provider key -- and each execution path silently depended on a different subset of +them. A `NeighborConnectivity` declaration produces all of them (ADR 0029), so what is left to +pin is that the *Python* name a declaration is reached through does not matter. """ import numpy as np @@ -42,34 +36,22 @@ class V(gtx.DimensionIndex): ... class E(gtx.DimensionIndex): ... -#: N1 == N3 == N4, but N2 differs: the tag is `TaggedOffDim.tag`, the variable is `off_a`. -class TaggedOffDim(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... - - -off_a = gtx.FieldOffset(TaggedOffDim.tag, source=E, target=(V, TaggedOffDim)) +class V2E(gtx.NeighborConnectivity[V, E]): + class Local(gtx.LocalDimensionIndex): ... -#: N1 == N2 == N4, but N3 differs: the local dimension is `Neigh`, the tag is `OffB`. -class Neigh(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +#: The declaration, reached through a different Python name. +off_a = V2E +#: Its local dimension, likewise. +Neigh = V2E.Local -OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) - - -def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases.Case: - """ - A `Case` whose offset provider holds exactly one connectivity, keyed on its tag. - - One entry per `Case` on purpose: DaCe walks every provider entry while building the - SDFG, and looks a connectivity up by its *local dimension's* name - (`gtir_to_sdfg.py`, constraint A4). A second, non-conforming entry would therefore - fail a program that does not even use it, and the cell under test would be measuring - the wrong thing. - """ +@pytest.fixture +def case(exec_alloc_descriptor) -> cases.Case: mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) # NOTE: `.asnumpy()`, not `.ndarray`: under a GPU allocator the latter is a device # array, and `simple_mesh` builds the table from NumPy anyway. - v2e_arr = mesh.offset_provider[cases_utils.V2EDim.tag].asnumpy() + v2e_arr = mesh.offset_provider[cases_utils.V2E].asnumpy() return cases.Case( ( None @@ -77,8 +59,8 @@ def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases else exec_alloc_descriptor ), offset_provider={ - tag: constructors.as_connectivity( - domain={V: v2e_arr.shape[0], local_dim: v2e_arr.shape[1]}, + off_a: constructors.as_connectivity( + domain={V: v2e_arr.shape[0], Neigh: v2e_arr.shape[1]}, codomain=E, data=v2e_arr, skip_value=None, @@ -91,83 +73,29 @@ def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases ) -@pytest.fixture -def case_tag_vs_variable_name(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, TaggedOffDim.tag, TaggedOffDim) - - -@pytest.fixture -def case_tag_vs_local_dim(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, "OffB", Neigh) - - -def _neighbor_table(case: cases.Case, tag: str) -> np.ndarray: - return case.offset_provider[tag].asnumpy() - +def _neighbor_table(case: cases.Case) -> np.ndarray: + return case.offset_provider[V2E].asnumpy() -# --- N2: the tag differs from the Python variable name ---------------------------- -# Lowering used to emit the *variable* name as the IR shift tag, so embedded and -# compiled execution of the same program needed different provider keys. - -def test_shift_tag_differs_from_variable_name(case_tag_vs_variable_name): +def test_shift_through_an_alias(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: return a(off_a[1]) - cases.verify_with_default_data( - case_tag_vs_variable_name, - foo, - lambda a: a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)[:, 1]], - ) - - -def test_reduction_tag_differs_from_variable_name(case_tag_vs_variable_name): - @gtx.field_operator - def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return neighbor_sum(a(off_a), axis=TaggedOffDim) - - cases.verify_with_default_data( - case_tag_vs_variable_name, - foo, - lambda a: np.sum(a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)], axis=1), - ) - - -# --- N3: the tag differs from the local dimension's name -------------------------- -# Lifted for the gtfn shift path by #1789; still required elsewhere, which is what -# the markers below record. - + cases.verify_with_default_data(case, foo, lambda a: a[_neighbor_table(case)[:, 1]]) -@pytest.mark.uses_offset_tag_differing_from_local_dim -def test_shift_tag_differs_from_local_dim_name(case_tag_vs_local_dim): - """ - Ensure a shift works with an offset tag that differs from the local dimension's name. - - If `NeighborConnectivityType.neighbor_dim` did not match the `FieldOffset` value, - gtfn would silently ignore the neighbor index, see - https://github.com/GridTools/gridtools/pull/1814. - """ +def test_reduction_through_an_alias(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return a(OffB[1]) + return neighbor_sum(a(off_a), axis=Neigh) - cases.verify_with_default_data( - case_tag_vs_local_dim, - foo, - lambda a: a[_neighbor_table(case_tag_vs_local_dim, "OffB")[:, 1]], - ) + cases.verify_with_default_data(case, foo, lambda a: np.sum(a[_neighbor_table(case)], axis=1)) -@pytest.mark.uses_offset_tag_differing_from_local_dim_in_reduction -def test_reduction_tag_differs_from_local_dim_name(case_tag_vs_local_dim): +def test_reduction_over_the_nested_name(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return neighbor_sum(a(OffB), axis=Neigh) + return neighbor_sum(a(V2E), axis=off_a.Local) - cases.verify_with_default_data( - case_tag_vs_local_dim, - foo, - lambda a: np.sum(a[_neighbor_table(case_tag_vs_local_dim, "OffB")], axis=1), - ) + cases.verify_with_default_data(case, foo, lambda a: np.sum(a[_neighbor_table(case)], axis=1)) diff --git a/tests/next_tests/toy_connectivity.py b/tests/next_tests/toy_connectivity.py index c368b04245..28fffdc732 100644 --- a/tests/next_tests/toy_connectivity.py +++ b/tests/next_tests/toy_connectivity.py @@ -21,22 +21,26 @@ class Edge(gtx.DimensionIndex): ... class Cell(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2V(gtx.NeighborConnectivity[Edge, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class V2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2V(gtx.NeighborConnectivity[Vertex, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) -V2V = gtx.FieldOffset(V2VDim.tag, source=Vertex, target=(Vertex, V2VDim)) +V2EDim = V2E.Local +E2VDim = E2V.Local +C2EDim = C2E.Local +V2VDim = V2V.Local # 3x3 periodic edges cells # 0 - 1 - 2 - 0 1 2 diff --git a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py index bef6a29c8a..1fa34616de 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py @@ -20,6 +20,7 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Field, DimensionIndex, @@ -49,10 +50,10 @@ class V(DimensionIndex): ... class E(DimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class V2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2V(LocalDimensionIndex): ... class C(DimensionIndex): ... @@ -61,7 +62,7 @@ class C(DimensionIndex): ... class K(DimensionIndex): ... -class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E2CO(LocalDimensionIndex): ... class A(DimensionIndex): ... @@ -76,7 +77,7 @@ class X(DimensionIndex): ... class Y(DimensionIndex): ... -class L(DimensionIndex, kind=DimensionKind.LOCAL): ... +class L(LocalDimensionIndex): ... class S(DimensionIndex): ... @@ -97,7 +98,7 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... class C2V(DimensionIndex): ... @@ -826,7 +827,6 @@ def test_premap_disjoint_inverse_image_raises(): def test_as_offset_1d(): # Dynamic per-point shift along I: out[i] == f[i + off[i]], full domain when all shifts in-bounds. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -834,7 +834,7 @@ def test_as_offset_1d(): off_arr = np.asarray([1, 0, -1, 0, 1, 0, -1, 0, 1, 0], dtype=int) off = common._field(off_arr, domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),))) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) assert np.all(result.ndarray == f.ndarray[np.arange(10) + off_arr]) @@ -842,7 +842,6 @@ def test_as_offset_1d(): def test_as_offset_narrow_offset_dtype_no_wrap(): # An int8 offset field over a domain larger than 128 must not wrap into the index table. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) N = 200 f = common._field( @@ -851,7 +850,7 @@ def test_as_offset_narrow_offset_dtype_no_wrap(): off_arr = np.zeros(N, dtype=np.int8) off = common._field(off_arr, domain=common.Domain(dims=(I,), ranges=(UnitRange(0, N),))) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(0, N),)) assert np.all(result.ndarray == f.ndarray) @@ -859,7 +858,6 @@ def test_as_offset_narrow_offset_dtype_no_wrap(): def test_as_offset_2d_shift_one_keep_other(): # Shift along I by a per-(i, j) offset, leave J: out[i, j] == f[i + off[i, j], j]. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) NI, NJ = 4, 3 dom = common.Domain(dims=(I, J), ranges=(UnitRange(0, NI), UnitRange(0, NJ))) @@ -867,7 +865,7 @@ def test_as_offset_2d_shift_one_keep_other(): off_arr = np.asarray([[1, 1, 0], [0, 0, 1], [1, -1, 0], [-1, 0, -1]], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == dom i = np.arange(NI)[:, None] @@ -877,7 +875,6 @@ def test_as_offset_2d_shift_one_keep_other(): def test_as_offset_boundary_narrows_domain(): # A uniform out-of-bounds shift narrows the result to the contiguous in-range sub-domain. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -886,7 +883,7 @@ def test_as_offset_boundary_narrows_domain(): np.full(10, -1, dtype=int), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) ) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(1, 10),)) assert np.all(result.ndarray == f.ndarray[0:9]) # out[i] == f[i - 1] @@ -894,7 +891,6 @@ def test_as_offset_boundary_narrows_domain(): def test_as_offset_scattered_oob_raises(): # An out-of-bounds shift in the interior cannot yield a contiguous domain. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -905,12 +901,11 @@ def test_as_offset_scattered_oob_raises(): ) with pytest.raises(ValueError, match="non-contiguous"): - f.premap(as_offset(Ioff, off)) + f.premap(as_offset(I, off)) def test_as_offset_introduces_dimension(): # `off` carries a dim the field lacks: the result gains it, out[i, j] == f[i + off[i, j]]. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -920,7 +915,7 @@ def test_as_offset_introduces_dimension(): off_arr, domain=common.Domain(dims=(I, J), ranges=(UnitRange(0, 10), UnitRange(0, 3))) ) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I, J), ranges=(UnitRange(0, 10), UnitRange(0, 3))) assert np.all(result.ndarray == f.ndarray[np.arange(10)[:, None] + off_arr]) @@ -928,14 +923,13 @@ def test_as_offset_introduces_dimension(): def test_as_offset_nonzero_origin(): # Field and offset over a domain that does not start at 0: indices must be shifted by the domain start. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) dom = common.Domain(dims=(I,), ranges=(UnitRange(2, 12),)) f = common._field(np.arange(10).astype(float), domain=dom) off_arr = np.asarray([1, 0, -1, 0, 1, 0, -1, 0, 1, 0], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == dom assert np.all(result.ndarray == f.ndarray[np.arange(10) + off_arr]) @@ -943,7 +937,6 @@ def test_as_offset_nonzero_origin(): def test_as_offset_2d_shift_second_axis(): # Shift along J (the non-leading axis) by a per-(i, j) offset, leave I: out[i, j] == f[i, j + off[i, j]]. - Joff = fbuiltins.FieldOffset("Joff", source=J, target=(J,)) NI, NJ = 3, 4 dom = common.Domain(dims=(I, J), ranges=(UnitRange(0, NI), UnitRange(0, NJ))) @@ -951,7 +944,7 @@ def test_as_offset_2d_shift_second_axis(): off_arr = np.asarray([[1, 1, 0, -1], [0, 0, 1, -1], [1, -1, 0, 0]], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Joff, off)) + result = f.premap(as_offset(J, off)) assert result.domain == dom i = np.arange(NI)[:, None] @@ -959,25 +952,13 @@ def test_as_offset_2d_shift_second_axis(): assert np.all(result.ndarray == f.ndarray[i, j + off_arr]) -def test_as_offset_non_cartesian_offset_raises(): - # `as_offset` only supports Cartesian (self-shift) offsets: single target equal to source. - - off_I = common._field( - np.zeros(3, dtype=int), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 3),)) - ) +def test_as_offset_local_dimension_raises(): + # `as_offset` shifts along a non-local dimension. off_V = common._field( np.zeros(3, dtype=int), domain=common.Domain(dims=(Vertex,), ranges=(UnitRange(0, 3),)) ) - - # 2-element target (neighbor offset) - V2E = fbuiltins.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) - with pytest.raises(ValueError, match="Cartesian"): - as_offset(V2E, off_V) - - # 1-element target but source != target[0] (cross-dim) - IfromJ = fbuiltins.FieldOffset("IfromJ", source=I, target=(J,)) - with pytest.raises(ValueError, match="Cartesian"): - as_offset(IfromJ, off_I) + with pytest.raises(ValueError, match="non-local dimension"): + as_offset(V2EDim, off_V) @pytest.mark.parametrize( diff --git a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py index 25284281ef..5efd6c9595 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py @@ -6,6 +6,8 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import typing + import pytest import gt4py.next as gtx @@ -21,33 +23,28 @@ class VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... class Dim(gtx.DimensionIndex): ... -class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class LocalDim(gtx.LocalDimensionIndex): ... -CartesianOffset = gtx.FieldOffset("CartesianOffset", source=Dim, target=(Dim,)) -UnstructuredOffset = gtx.FieldOffset(LocalDim.tag, source=Dim, target=(Dim, LocalDim)) +class UnstructuredOffset(gtx.NeighborConnectivity[Dim, Dim]): + Local: typing.TypeAlias = LocalDim def test_domain_deduction_cartesian(): - assert _deduce_grid_type(None, {CartesianOffset}) == gtx.GridType.CARTESIAN assert _deduce_grid_type(None, {Dim}) == gtx.GridType.CARTESIAN + assert _deduce_grid_type(None, {HDim, VDim}) == gtx.GridType.CARTESIAN def test_domain_deduction_unstructured(): assert _deduce_grid_type(None, {UnstructuredOffset}) == gtx.GridType.UNSTRUCTURED assert _deduce_grid_type(None, {LocalDim}) == gtx.GridType.UNSTRUCTURED - # source and target share `.value` but differ in `.kind` -> not Cartesian - CrossKindOffset = gtx.FieldOffset("CrossKind", source=HDim, target=(VDim,)) - assert _deduce_grid_type(None, {CrossKindOffset}) == gtx.GridType.UNSTRUCTURED - # LOCAL self-loop is unstructured - LocalSelfOffset = gtx.FieldOffset("LocalSelf", source=LocalDim, target=(LocalDim,)) - assert _deduce_grid_type(None, {LocalSelfOffset}) == gtx.GridType.UNSTRUCTURED def test_domain_complies_with_request_cartesian(): - assert _deduce_grid_type(gtx.GridType.CARTESIAN, {CartesianOffset}) == gtx.GridType.CARTESIAN - with pytest.raises(ValueError, match="unstructured.*FieldOffset.*found"): + assert _deduce_grid_type(gtx.GridType.CARTESIAN, {Dim}) == gtx.GridType.CARTESIAN + with pytest.raises(ValueError, match="NeighborConnectivity.*local dimension was found"): _deduce_grid_type(gtx.GridType.CARTESIAN, {UnstructuredOffset}) + with pytest.raises(ValueError, match="NeighborConnectivity.*local dimension was found"): _deduce_grid_type(gtx.GridType.CARTESIAN, {LocalDim}) @@ -57,6 +54,4 @@ def test_domain_complies_with_request_unstructured(): == gtx.GridType.UNSTRUCTURED ) # unstructured is ok, even if we don't have unstructured offsets - assert ( - _deduce_grid_type(gtx.GridType.UNSTRUCTURED, {CartesianOffset}) == gtx.GridType.UNSTRUCTURED - ) + assert _deduce_grid_type(gtx.GridType.UNSTRUCTURED, {Dim}) == gtx.GridType.UNSTRUCTURED diff --git a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py index 337c7ce67f..1fc3ca6782 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py @@ -30,8 +30,6 @@ class IDim(gtx.DimensionIndex): ... -IOff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) - # A PEP 695 alias whose value raises when it is evaluated, standing in for the # common case of a typo'd dtype ('np.foat64') inside an alias definition. _empty_module = types.ModuleType("_empty_module") @@ -361,20 +359,6 @@ def broken(a: BrokenFieldAlias) -> gtx.Field[[IDim], float64]: assert re.search(r"\| +\^{19}(?!\^)", str(err)), str(err) -def test_unindexed_cartesian_offset_names_the_offset_as_written(): - # The tag of 'IOff' is 'Ioff'; the message has to quote what the user wrote. - def unindexed(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: - return a(IOff) - - err = parse_error(unindexed) - - assert err.message == "Cannot shift by the Cartesian offset 'IOff' without an index." - assert err.hints == ["Give the displacement, e.g. 'IOff[1]'."] - rendered = str(err) - assert "return a(IOff)" in rendered - assert re.search(r"\| +\^{4}(?!\^)", rendered), rendered - - def test_indexed_dimension_shift_is_rejected_with_a_hint(): def indexed_shift(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: return a((IDim + 1)[0]) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index f696daf1b4..6d423a52eb 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py @@ -49,24 +49,19 @@ class Edge(gtx.DimensionIndex): ... class Vertex(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) +V2EDim = V2E.Local class TDim(gtx.DimensionIndex): ... -TOff = gtx.FieldOffset("TDim", source=TDim, target=(TDim,)) - - -#: An offset whose tag differs from the name of the Python variable it is bound to, and -#: from the name of its local dimension. Lowering must emit the *tag*. -class RenamedV2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... - - -renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) +#: A connectivity reached through a name other than its declaration's. Lowering must emit the +#: local dimension's tag, not the variable name. +renamed_v2e = V2E class UDim(gtx.DimensionIndex): ... @@ -166,7 +161,7 @@ def foo_float(inp: gtx.Field[[TDim], float64]): def test_as_offset(): def foo(inp: gtx.Field[[TDim], float64], offset: gtx.Field[[TDim], int]): - return inp(as_offset(TOff, offset)) + return inp(as_offset(TDim, offset)) parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) @@ -802,7 +797,7 @@ def foo(edge_f: gtx.Field[gtx.Dims[Edge], float64]): def test_unstructured_shift_lowering_emits_offset_tag_not_variable_name(): - """The IR shift tag is the offset's tag, not the variable the offset is bound to.""" + """The IR shift tag is the local dimension's tag, not the variable it is reached through.""" def foo(edge_f: gtx.Field[[Edge], float64]): return edge_f(renamed_v2e[1]) @@ -810,7 +805,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) - reference = im.as_fieldop(im.lambda_("__it")(im.deref(im.shift("RenamedTag", 1)("__it"))))( + reference = im.as_fieldop(im.lambda_("__it")(im.deref(im.shift(V2EDim.tag, 1)("__it"))))( "edge_f" ) @@ -824,7 +819,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) - reference = im.as_fieldop_neighbors("RenamedTag", "edge_f") + reference = im.as_fieldop_neighbors(V2EDim.tag, "edge_f") assert lowered.expr == reference @@ -985,7 +980,7 @@ def foo(inp: gtx.Field[[TDim], float64]): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( im.ref("inp"), - im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (TDim, UDim))), ) @@ -999,7 +994,7 @@ def foo(): assert lowered.id == "foo" assert lowered.expr == im.call("broadcast")( 1, - im.make_tuple(*(itir.AxisLiteral(value=dim.tag, kind=dim.kind) for dim in (TDim, UDim))), + im.make_tuple(*(itir.AxisLiteral(value=dim.tag) for dim in (TDim, UDim))), ) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py index 290d2914fc..7051400ed8 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py @@ -19,7 +19,7 @@ """ import gt4py.next as gtx -from gt4py.next import Dims, Dimension, DimensionIndex, float64, neighbor_sum +from gt4py.next import Dims, Dimension, DimensionIndex, LocalDimensionIndex, float64, neighbor_sum from gt4py.next.ffront import source_utils from gt4py.next.ffront.source_utils import get_closure_vars_from_function @@ -30,10 +30,11 @@ class Cell(DimensionIndex): ... class Edge(DimensionIndex): ... -class C2EDim(DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2EDim = C2E.Local CField = gtx.Field[Dims[Cell], float64] EField = gtx.Field[Dims[Edge], float64] diff --git a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py index 88a8c640d8..881372c568 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py @@ -18,8 +18,9 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, Field, - FieldOffset, astype, broadcast, errors, @@ -50,7 +51,11 @@ class X(DimensionIndex): ... class Y(DimensionIndex): ... -class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class Y2X(NeighborConnectivity[Y, X]): + class Local(LocalDimensionIndex): ... + + +Y2XDim = Y2X.Local class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... @@ -71,7 +76,11 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + + +V2EDim = V2E.Local class IDim(DimensionIndex): ... @@ -286,7 +295,6 @@ def domain_comparison(a: Field[[TDim], float], b: Field[[TDim], float]): @pytest.fixture def premap_setup(): - Y2X = FieldOffset(Y2XDim.tag, source=X, target=(Y, Y2XDim)) return X, Y, Y2XDim, Y2X @@ -547,41 +555,33 @@ def return_undefined(): def test_as_offset_dim(): - Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) - def as_offset_dim(a: Field[[ADim, BDim], float], b: Field[[ADim], int]): - return a(as_offset(Boff, b)) + return a(as_offset(BDim, b)) with pytest.raises(errors.DSLError, match=f"not in list of offset field dimensions"): _ = FieldOperatorParser.apply_to_function(as_offset_dim) def test_as_offset_dtype(): - Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) - def as_offset_dtype(a: Field[[ADim, BDim], float], b: Field[[BDim], float]): - return a(as_offset(Boff, b)) + return a(as_offset(BDim, b)) with pytest.raises(errors.DSLError, match=f"expected integer for offset field dtype"): _ = FieldOperatorParser.apply_to_function(as_offset_dtype) -def test_as_offset_non_cartesian(): - V2E = FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) - +def test_as_offset_non_dimension(): def as_offset_neighbor(a: Field[[Edge], float], b: Field[[Edge], int]): return a(as_offset(V2E, b)) - with pytest.raises(errors.DSLError, match="Cartesian"): + with pytest.raises(errors.DSLError, match="Expected 1st argument to be of type"): _ = FieldOperatorParser.apply_to_function(as_offset_neighbor) - IfromJ = FieldOffset("IfromJ", source=IDim, target=(JDim,)) - - def as_offset_cross_dim(a: Field[[IDim], float], b: Field[[IDim], int]): - return a(as_offset(IfromJ, b)) + def as_offset_local_dim(a: Field[[Vertex, V2EDim], float], b: Field[[Vertex, V2EDim], int]): + return a(as_offset(V2EDim, b)) - with pytest.raises(errors.DSLError, match="Cartesian"): - _ = FieldOperatorParser.apply_to_function(as_offset_cross_dim) + with pytest.raises(errors.DSLError, match="non-local dimension"): + _ = FieldOperatorParser.apply_to_function(as_offset_local_dim) vpfloat: TypeAlias = float32 diff --git a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py index 7120882c75..793f802c9f 100644 --- a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py +++ b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py @@ -33,13 +33,13 @@ class Vertex(common.DimensionIndex): ... class Edge(common.DimensionIndex): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... -class V2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2VDim(common.LocalDimensionIndex): ... a_range = domain_utils.SymbolicRange(0, 10) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py index 12c0b75649..842e84fc4e 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py @@ -10,6 +10,7 @@ import pytest import gt4py.next as gtx +from gt4py.next import common from gt4py.next.embedded import context as embedded_context from gt4py.next.iterator import embedded, runtime from gt4py.next.iterator.builtins import ( @@ -29,10 +30,11 @@ class E(gtx.DimensionIndex): ... class V(gtx.DimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2V(gtx.NeighborConnectivity[E, V]): + class Local(gtx.LocalDimensionIndex): ... -E2V = gtx.FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) +E2VDim = E2V.Local # 0 --0-- 1 --1-- 2 @@ -68,7 +70,7 @@ def testee(): ref = np.asarray([[42.0], [42.0]]) assert result.domain.dims[0] == E - assert result.domain.dims[1] == embedded._CONST_DIM # this is implementation detail + assert result.domain.dims[1] == common.ConstList # this is implementation detail assert result.shape[1] == 1 # this is implementation detail np.testing.assert_array_equal(result.asnumpy(), ref) @@ -151,6 +153,6 @@ def testee(): ref = np.asarray([[43.0], [43.0]]) assert result.domain.dims[0] == E - assert result.domain.dims[1] == embedded._CONST_DIM # this is implementation detail + assert result.domain.dims[1] == common.ConstList # this is implementation detail assert result.shape[1] == 1 # this is implementation detail np.testing.assert_array_equal(result.asnumpy(), ref) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py b/tests/next_tests/unit_tests/iterator_tests/test_pretty_parser.py index 83a16a980c..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..c02336e6e0 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_pretty_printer.py @@ -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.DimensionIndex): ... -def test_axis_literal_vertical(): - testee = ir.AxisLiteral(value="I", kind=ir.DimensionKind.VERTICAL) - expected = "Iᵥ" - actual = pformat(testee) - assert actual == expected +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class LocalDim(gtx.LocalDimensionIndex): ... + + +@pytest.mark.parametrize("dim, suffix", [(IDim, "ₕ"), (KDim, "ᵥ"), (LocalDim, "ₗ")]) +def test_axis_literal(dim, suffix): + # the suffix is the resolved dimension's kind + assert pformat(ir.AxisLiteral(value=dim.tag)) == f"{dim.tag}{suffix}" + + +def test_axis_literal_of_unresolvable_tag(): + # 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_type_inference.py b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py index 94f0610c06..4f580b110c 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py @@ -101,9 +101,7 @@ def expression_test_cases(): bool_type, ), ( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), it_ts.NamedRangeType(dim=Vertex), ), ( @@ -112,9 +110,7 @@ def expression_test_cases(): ), ( im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ) + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1) ), ts.DomainType(dims=[Vertex]), ), @@ -443,10 +439,8 @@ def test_cartesian_fencil_definition(): def test_unstructured_fencil_definition(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), - im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( @@ -510,10 +504,8 @@ def test_function_definition(): def test_fencil_with_nb_field_input(): mesh = simple_mesh(None) unstructured_domain = im.call("unstructured_domain")( - im.named_range( - itir.AxisLiteral(value=Vertex.tag, kind=common.DimensionKind.HORIZONTAL), 0, 1 - ), - im.named_range(itir.AxisLiteral(value=KDim.tag, kind=common.DimensionKind.VERTICAL), 0, 1), + im.named_range(itir.AxisLiteral(value=Vertex.tag), 0, 1), + im.named_range(itir.AxisLiteral(value=KDim.tag), 0, 1), ) testee = itir.Program( diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py index eafe45b9d3..23c4dbaf07 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py @@ -45,7 +45,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... float_i_field = ts.FieldType(dims=[IDim], dtype=float_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py index 5bceba0e53..7dfbf60934 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py @@ -17,7 +17,7 @@ from gt4py.next.type_system import type_specifications as ts -class Neighbor(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Neighbor(common.LocalDimensionIndex): ... class IDim(common.DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py index 3f6cd212ae..d58b728e10 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py @@ -25,7 +25,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py index 5451a2dc44..7c81aa5514 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py @@ -11,7 +11,7 @@ from gt4py.next import common, utils from gt4py.next.iterator import ir from gt4py.next.iterator.ir_utils import ir_makers as im -from gt4py.next.iterator.transforms.unroll_reduce import UnrollReduce, _get_partial_offset_tags +from gt4py.next.iterator.transforms.unroll_reduce import UnrollReduce, _get_partial_local_dims from gt4py.next.type_system import type_specifications as ts @@ -27,10 +27,10 @@ class dummy_neighbor(common.DimensionIndex): ... #: The local dimensions of the neighbor lists under test. Each one's `tag` is also its IR offset #: string and its offset-provider key: `UnrollReduce` looks a connectivity up by the local #: dimension of the list it reduces, so those three names must be a single string (ADR 0028). -class Dim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim(common.LocalDimensionIndex): ... -class Dim2(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim2(common.LocalDimensionIndex): ... def dummy_connectivity_type(max_neighbors: int, has_skip_values: bool): @@ -98,10 +98,10 @@ def reduction_if(): "reduction_if", ], ) -def test_get_partial_offsets(reduction, request): - partial_offsets = _get_partial_offset_tags(request.getfixturevalue(reduction).args) +def test_get_partial_local_dims(reduction, request): + partial_local_dims = _get_partial_local_dims(request.getfixturevalue(reduction).args) - assert set(partial_offsets) == {Dim.tag} + assert set(partial_local_dims) == {Dim} def _expected(red, max_neighbors, has_skip_values, shifted_arg=0): diff --git a/tests/next_tests/unit_tests/otf_tests/test_runners.py b/tests/next_tests/unit_tests/otf_tests/test_runners.py index 2ea966ff7e..f75471fcda 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_runners.py +++ b/tests/next_tests/unit_tests/otf_tests/test_runners.py @@ -31,7 +31,7 @@ class Vertex(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... @pytest.fixture diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py index faea3c7764..1a02371357 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py @@ -208,7 +208,7 @@ def verify_testee(): def test_dace_fastcall_with_connectivity(unstructured_case, monkeypatch): """Test reuse of SDFG arguments between program calls by means of SDFG fastcall API.""" - connectivity_E2V = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + connectivity_E2V = unstructured_case.offset_provider[E2V].asnumpy() @gtx.field_operator def testee(a: cases.VField) -> cases.EField: diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index dd58299c02..2d9c31ca1c 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -375,7 +375,7 @@ def testee(a: cases.VField, b: cases.VField): ), ) - SIMPLE_MESH = cases_utils.simple_mesh(None) + SIMPLE_MESH = cases_utils.ir_level(cases_utils.simple_mesh(None)) offset_provider = SIMPLE_MESH.offset_provider test_case = cases.Case.from_mesh_descriptor(SIMPLE_MESH, backend=backend, allocator=backend) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index 1bec200fad..467b0b6de7 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -29,6 +29,7 @@ ) from gt4py.next.type_system import type_specifications as ts +from next_tests.integration_tests import cases_utils from next_tests.integration_tests.cases_utils import ( V2E, Edge, @@ -80,7 +81,7 @@ def _translate_gtir_to_sdfg( @pytest.mark.parametrize("has_unit_stride", [False, True]) @pytest.mark.parametrize("disable_field_origin", [False, True]) def test_find_constant_symbols(has_unit_stride, disable_field_origin): - SKIP_VALUE_MESH = skip_value_mesh(None) + SKIP_VALUE_MESH = cases_utils.ir_level(skip_value_mesh(None)) ir = itir.Program( id="find_constant_symbols_sdfg", @@ -94,7 +95,7 @@ def test_find_constant_symbols(has_unit_stride, disable_field_origin): itir.SetAt( expr=im.as_fieldop( im.lambda_("it")(im.reduce("plus", im.literal_from_value(1.0))(im.deref("it"))) - )(im.as_fieldop_neighbors(V2E.value, "x")), + )(im.as_fieldop_neighbors(V2E.Local.tag, "x")), domain=im.get_field_domain(gtx_common.GridType.UNSTRUCTURED, "y", VFTYPE.dims), target=itir.SymRef(id="y"), ) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py index 6fdec5a21f..78f1045eb9 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py @@ -18,3 +18,34 @@ def test_safe_replace_symbolic(): assert gtir_to_sdfg_utils.safe_replace_symbolic( dace.symbolic.pystr_to_symbolic("x*x + y"), symbol_mapping={"x": "y", "y": "x"} ) == dace.symbolic.pystr_to_symbolic("y*y + x") + + +def test_local_dimension_size(): + import numpy as np + + from gt4py._core import definitions as core_defs + from gt4py.next import common + from gt4py.next.program_processors.runners.dace import sdfg_args + + from next_tests.toy_connectivity import V2E, V2EDim, Vertex, Edge + + def conn_type(max_neighbors: int) -> common.NeighborConnectivityType: + return common.NeighborConnectivityType( + domain=(Vertex, V2EDim), + codomain=Edge, + skip_value=None, + dtype=core_defs.dtype(np.int32), + max_neighbors=max_neighbors, + ) + + sharer_tag = "some.module.V2EShared" + table_types = {V2E.offset_tag: conn_type(4), sharer_tag: conn_type(4)} + # a field finds the size in the table keyed by the local dimension + assert sdfg_args.local_dimension_size("a_field", V2EDim, table_types) == 4 + # a connectivity array has its own + conn_array = sdfg_args.connectivity_identifier(sharer_tag) + assert sdfg_args.local_dimension_size(conn_array, V2EDim, {sharer_tag: conn_type(4)}) == 4 + # a field over a local dimension bound only through a sharing connectivity + assert sdfg_args.local_dimension_size("a_field", V2EDim, {sharer_tag: conn_type(4)}) == 4 + with pytest.raises(KeyError, match="No connectivity over the local dimension"): + sdfg_args.local_dimension_size("a_field", V2EDim, {}) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py index a1fb215882..7a9e2e6114 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py @@ -26,6 +26,7 @@ from gt4py.next.iterator.transforms import pass_manager from gt4py.next.type_system import type_specifications as ts +from next_tests.integration_tests import cases_utils from next_tests.integration_tests.cases_utils import ( E2VDim, C2VDim, @@ -72,8 +73,8 @@ def allow_view_arguments(): IOff = im.cartesian_offset(IDim, IDim) # Cartesian shifts are self-describing (`CartesianOffset`), so no offset provider entry is needed. CARTESIAN_OFFSETS: dict = {} -SIMPLE_MESH: MeshDescriptor = simple_mesh(None) -SKIP_VALUE_MESH: MeshDescriptor = skip_value_mesh(None) +SIMPLE_MESH: MeshDescriptor = cases_utils.ir_level(simple_mesh(None)) +SKIP_VALUE_MESH: MeshDescriptor = cases_utils.ir_level(skip_value_mesh(None)) SIZE_TYPE = ts.ScalarType(ts.ScalarKind.INT32) FSYMBOLS = dict( **{gtx_dace_args.range_start_symbol("w", IDim).name: 0}, diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index 500a355dbb..6629711027 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -22,6 +22,7 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Infinity, UnitRange, @@ -57,19 +58,19 @@ class I(common.DimensionIndex): ... class I_half(common.DimensionIndex): ... -class C2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E(LocalDimensionIndex): ... -class V2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(LocalDimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C(LocalDimensionIndex): ... -class E2C2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2V(LocalDimensionIndex): ... class ECDim(DimensionIndex): ... @@ -917,3 +918,12 @@ def test_a_dimension_cannot_be_staggered_twice(self): def test_resolve_rejects_a_bracketed_tag_of_another_owner(self): with pytest.raises(ValueError, match="not a parametrized dimension"): common.resolve(f"{KDim.tag}[{KDim.tag}]") + + +def test_resolve_loaded(): + # `resolve_loaded` never imports: it answers for what is loaded and gives up otherwise + assert common.resolve_loaded(IDim.tag) is IDim + assert common.resolve_loaded(common.Staggered[IDim].tag) is common.Staggered[IDim] + assert common.resolve_loaded("not_imported_anywhere.IDim") is None + assert common.resolve_loaded(f"{__name__}.does_not_exist") is None + assert common.resolve_loaded(f"{__name__}.test_resolve_loaded") is None # not a dimension diff --git a/tests/next_tests/unit_tests/test_custom_layout_allocators.py b/tests/next_tests/unit_tests/test_custom_layout_allocators.py index f5257506fb..53312af3ab 100644 --- a/tests/next_tests/unit_tests/test_custom_layout_allocators.py +++ b/tests/next_tests/unit_tests/test_custom_layout_allocators.py @@ -29,13 +29,13 @@ class D2(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class D0_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... -class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D1_local(common.LocalDimensionIndex): ... class D2_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... -class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D2_local(common.LocalDimensionIndex): ... class D1_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py new file mode 100644 index 0000000000..7462d38d2b --- /dev/null +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -0,0 +1,654 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import pickle +import textwrap +import typing + +import numpy as np +import pytest + +from gt4py._core import definitions as core_defs +from gt4py.next import common +from gt4py.next.common import ( + DimensionIndex, + DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, + NeighborConnectivityType, +) +from gt4py.next.ffront import transform_utils +from gt4py.next.type_system import type_specifications as ts, type_translation + + +class Vertex(DimensionIndex): ... + + +class Edge(DimensionIndex): ... + + +class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... + + +class V2E(NeighborConnectivity[Vertex, Edge], max_neighbors=4, min_neighbors=3): + class Local(LocalDimensionIndex): ... + + +class E2V(NeighborConnectivity[Edge, Vertex]): + class Local(LocalDimensionIndex): ... + + +class LsqCoeff(LocalDimensionIndex, size=3): ... + + +class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local = V2E.Local + + +def _declare(source: str) -> dict: + """ + Run `source` as the body of a throwaway module. + + Declarations have to be at module level, so error cases cannot simply be written inside + the test function: the `` check would fire before the one under test. + """ + namespace = { + "__name__": __name__, + "typing": typing, + "DimensionIndex": DimensionIndex, + "DimensionKind": DimensionKind, + "LocalDimensionIndex": LocalDimensionIndex, + "NeighborConnectivity": NeighborConnectivity, + "Vertex": Vertex, + "Edge": Edge, + "KDim": KDim, + "V2E": V2E, + "ConstList": common.ConstList, + } + exec(textwrap.dedent(source), namespace) + return namespace + + +class TestDeclaration: + def test_owner_and_dimensions(self): + assert V2E.Local.owner is V2E + assert V2E.origin is Vertex + assert V2E.codomain is Edge + assert V2E.Local.kind is DimensionKind.LOCAL + assert issubclass(V2E.Local, DimensionIndex) + + def test_counts(self): + assert (V2E.Local.max_neighbors, V2E.Local.min_neighbors) == (4, 3) + assert (E2V.Local.max_neighbors, E2V.Local.min_neighbors) == (None, None) + + def test_ownerless_local(self): + assert LsqCoeff.owner is None + assert (LsqCoeff.max_neighbors, LsqCoeff.min_neighbors) == (3, 3) + assert LsqCoeff.kind is DimensionKind.LOCAL + + def test_counts_from_local_size(self): + ns = _declare( + """ + class C2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex, size=3): ... + """ + ) + assert (ns["C2E"].Local.max_neighbors, ns["C2E"].Local.min_neighbors) == (3, 3) + + def test_identity(self): + assert V2E.tag == f"{__name__}.V2E" + assert V2E.Local.tag == f"{__name__}.V2E.Local" + assert common.resolve(V2E.Local.tag) is V2E.Local + assert str(V2E) == "V2E" + assert repr(V2E) == V2E.tag + + def test_pickle_by_reference(self): + assert pickle.loads(pickle.dumps(V2E)) is V2E + assert pickle.loads(pickle.dumps(V2E.Local)) is V2E.Local + + def test_hashable(self): + assert {V2E: 1}[V2E] == 1 + + def test_type_parameter_subscription(self): + alias = NeighborConnectivity[Vertex, Edge] + assert typing.get_origin(alias) is NeighborConnectivity + assert typing.get_args(alias) == (Vertex, Edge) + + def test_bool_is_not_a_neighbor_index(self): + with pytest.raises(TypeError): + V2E[True] + + def test_not_instantiable(self): + with pytest.raises(TypeError, match="cannot be instantiated"): + V2E() + + +class TestDeclarationErrors: + @pytest.mark.parametrize( + "source, match", + [ + ( + """ + class C(NeighborConnectivity[Vertex, Edge]): ... + """, + "must declare its local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge]): + class Local(DimensionIndex): ... + """, + "must declare its local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=5): + Local: typing.TypeAlias = V2E.Local + """, + "contradicts the local dimension it shares with 'V2E'", + ), + ( + """ + class C(NeighborConnectivity[Edge, Edge]): + Local = V2E.Local + """, + "cannot share the local dimension of 'V2E'", + ), + ( + """ + class C(NeighborConnectivity): + class Local(LocalDimensionIndex): ... + """, + "must derive from 'NeighborConnectivity\\[Origin, Codomain\\]'", + ), + ( + """ + class C(V2E): + class Local(LocalDimensionIndex): ... + """, + "must derive from 'NeighborConnectivity\\[Origin, Codomain\\]'", + ), + ( + """ + class C(NeighborConnectivity[V2E.Local, Edge]): + class Local(LocalDimensionIndex): ... + """, + "'Origin' must be a non-local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, int]): + class Local(LocalDimensionIndex): ... + """, + "'Codomain' must be a non-local dimension", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=2, min_neighbors=3): + class Local(LocalDimensionIndex): ... + """, + "exceeds 'max_neighbors'", + ), + ( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=4): + class Local(LocalDimensionIndex, size=3): ... + """, + "contradicts the size", + ), + ( + """ + class L(LocalDimensionIndex, kind=DimensionKind.HORIZONTAL): ... + """, + "cannot have kind", + ), + ( + """ + class L(LocalDimensionIndex, size=1.5): ... + """, + "must be an integer", + ), + ( + """ + class L(DimensionIndex, kind=DimensionKind.LOCAL): ... + """, + "subclassing 'LocalDimensionIndex'", + ), + ], + ) + def test_rejected(self, source, match): + with pytest.raises(TypeError, match=match): + _declare(source) + + def test_negative_count(self): + with pytest.raises(ValueError, match="non-negative"): + _declare( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=-1): + class Local(LocalDimensionIndex): ... + """ + ) + + def test_adopting_an_ownerless_local(self): + ns = _declare( + """ + class Coeff(LocalDimensionIndex, size=3): ... + + class C(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = Coeff + """ + ) + assert ns["Coeff"].owner is ns["C"] + assert (ns["Coeff"].max_neighbors, ns["Coeff"].min_neighbors) == (3, 3) + + def test_sharing_a_local_dimension(self): + ns = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + ) + shared = ns["V2EShared"] + assert shared.Local is V2E.Local + assert V2E.Local.owner is V2E + assert V2E.offset_tag == V2E.Local.tag + assert shared.offset_tag == shared.tag + assert shared.__gt_type__().tag == shared.tag + + def test_sharing_with_consistent_counts(self): + ns = _declare( + """ + class V2EShared4(NeighborConnectivity[Vertex, Edge], max_neighbors=4): + Local = V2E.Local + """ + ) + assert ns["V2EShared4"].Local is V2E.Local + + def test_non_integer_index(self): + with pytest.raises(TypeError, match="indexed by an integer"): + V2E[Vertex] + + def test_base_is_not_a_declaration(self): + with pytest.raises(TypeError, match="not a connectivity declaration"): + NeighborConnectivity.__gt_type__() + + def test_function_local_declaration(self): + with pytest.raises(TypeError, match="module level"): + + class C(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = LsqCoeff + + def test_subclass_of_owned_local_is_ownerless(self): + ns = _declare( + """ + class Other(V2E.Local): ... + """ + ) + assert ns["Other"].owner is None + + def test_local_dimension_cannot_be_staggered(self): + with pytest.raises(TypeError, match="cannot be staggered"): + common.Staggered[V2E.Local] + + +def _table_type( + domain=(Vertex, V2E.Local), + codomain=Edge, + max_neighbors=4, + skip_value=common._DEFAULT_SKIP_VALUE, + dtype=np.int32, +) -> NeighborConnectivityType: + return NeighborConnectivityType( + domain=domain, + codomain=codomain, + skip_value=skip_value, + dtype=core_defs.dtype(dtype), + max_neighbors=max_neighbors, + ) + + +class TestCheckNeighborTable: + def test_matching_type(self): + common.check_neighbor_table(V2E, _table_type()) + + def test_matching_table(self): + from gt4py.next import constructors + + table = constructors.as_connectivity( + domain={Edge: 2, E2V.Local: 2}, codomain=Vertex, data=np.array([[0, 1], [1, 2]]) + ) + common.check_neighbor_table(E2V, table) + + def test_undeclared_counts_accept_any_table(self): + common.check_neighbor_table( + E2V, _table_type(domain=(Edge, E2V.Local), codomain=Vertex, max_neighbors=7) + ) + + @pytest.mark.parametrize( + "kwargs, match", + [ + ({"domain": (Vertex, E2V.Local)}, "its domain is"), + ({"domain": (Edge, V2E.Local)}, "its domain is"), + ({"codomain": Vertex}, "its codomain is"), + ({"dtype": np.float64}, "is not integral"), + ({"max_neighbors": 5}, "expected max_neighbors=4"), + ({"skip_value": None}, "requires a skip value"), + ], + ) + def test_mismatch(self, kwargs, match): + with pytest.raises(ValueError, match=match): + common.check_neighbor_table(V2E, _table_type(**kwargs)) + + def test_min_neighbors_exceeds_table(self): + ns = _declare( + """ + class MinOnly(NeighborConnectivity[Vertex, Edge], min_neighbors=5): + class Local(LocalDimensionIndex): ... + """ + ) + min_only = ns["MinOnly"] + for skip_value in (None, common._DEFAULT_SKIP_VALUE): + with pytest.raises(ValueError, match="min_neighbors=5 exceeds"): + common.check_neighbor_table( + min_only, + _table_type( + domain=(Vertex, min_only.Local), max_neighbors=3, skip_value=skip_value + ), + ) + + def test_bool_table_is_not_integral(self): + with pytest.raises(ValueError, match="is not integral"): + common.check_neighbor_table(V2E, _table_type(dtype=bool)) + + def test_not_a_neighbor_table(self): + with pytest.raises(ValueError, match="expected a neighbor table"): + common.check_neighbor_table(V2E, common.CartesianConnectivity(Vertex, 1)) + + def test_skip_value_without_missing_neighbors(self): + ns = _declare( + """ + class Full(NeighborConnectivity[Vertex, Edge], max_neighbors=2, min_neighbors=2): + class Local(LocalDimensionIndex): ... + """ + ) + full = ns["Full"] + with pytest.raises(ValueError, match="has skip value"): + common.check_neighbor_table( + full, _table_type(domain=(Vertex, full.Local), max_neighbors=2) + ) + + +class TestFrontendIntegration: + def test_from_value_is_an_offset(self): + # NOTE: pins the `__gt_type__` branch of `from_value` ahead of the dimension branch; a + # connectivity declaration is a class, like a dimension. + assert type_translation.from_value(V2E) == ts.OffsetType( + source=Edge, target=(Vertex, V2E.Local), tag=V2E.Local.tag + ) + + def test_neighbor_index_accepts_numpy_integers(self): + from gt4py.next import constructors, embedded + + table = constructors.as_connectivity( + domain={Vertex: 2, V2E.Local: 4}, + codomain=Edge, + data=np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), + ) + with embedded.context.update(offset_provider={V2E: table}): + assert np.array_equal(V2E[np.int32(1)].asnumpy(), V2E[1].asnumpy()) + + def test_attribute_errors_are_dsl_errors(self): + from gt4py.next import errors, field_operator + from gt4py.next.ffront.func_to_foast import FieldOperatorParser + from gt4py.next import Dims, Field + + def origin_of(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return a(V2E.origin) + + with pytest.raises(errors.DSLError, match="has no attribute 'origin'"): + FieldOperatorParser.apply_to_function(origin_of) + + def test_fingerprint_covers_the_declaration(self): + from gt4py.next import fingerprinting + + def fingerprint_of(source: str) -> str: + # lenient: `_declare` classes are not importable, as in a re-run notebook cell + return fingerprinting.lenient_fingerprinter(_declare(source)["C"]) + + base = fingerprint_of( + """ + class C(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + """ + ) + swapped = fingerprint_of( + """ + class C(NeighborConnectivity[Edge, Vertex]): + class Local(LocalDimensionIndex): ... + """ + ) + counted = fingerprint_of( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=3): + class Local(LocalDimensionIndex): ... + """ + ) + assert len({base, swapped, counted}) == 3 + assert fingerprinting.strict_fingerprinter(V2E) != fingerprinting.strict_fingerprinter(E2V) + + def test_grid_type_deduction(self): + assert ( + transform_utils._deduce_grid_type(None, [Vertex, V2E]) is common.GridType.UNSTRUCTURED + ) + with pytest.raises(ValueError, match="CARTESIAN"): + transform_utils._deduce_grid_type(common.GridType.CARTESIAN, [V2E]) + + +def _table(domain=(Vertex, V2E.Local), codomain=Edge, data=((0, 1, 2, 3), (1, 2, 3, 0))): + from gt4py.next import constructors + + data = np.array(data) + return constructors.as_connectivity( + domain=dict(zip(domain, data.shape)), + codomain=codomain, + data=data, + skip_value=common._DEFAULT_SKIP_VALUE, + ) + + +def test_redefined_declaration_with_an_adopted_local(monkeypatch): + """Re-running a cell must re-own the adopted local dimension, not become a sharer.""" + import sys + import types as pytypes + + module = pytypes.ModuleType("_readopted_connectivity_module") + monkeypatch.setitem(sys.modules, module.__name__, module) + source = textwrap.dedent( + """ + import typing + + from gt4py.next.common import DimensionIndex, LocalDimensionIndex, NeighborConnectivity + + class V(DimensionIndex): ... + class E(DimensionIndex): ... + class V2EDim(LocalDimensionIndex, size={n}): ... + class V2E(NeighborConnectivity[V, E], max_neighbors={n}): + Local: typing.TypeAlias = V2EDim + """ + ) + exec(source.format(n=4), module.__dict__) + assert module.V2E.offset_tag == module.V2EDim.tag + + # the redefinition takes ownership over again, and its counts are checked against `size=` + exec(source.format(n=2), module.__dict__) + assert module.V2EDim.owner is module.V2E + assert module.V2E.offset_tag == module.V2EDim.tag + assert module.V2EDim.max_neighbors == 2 + + +def test_local_dimension_of(): + shared = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2EShared"] + assert common.local_dimension_of(V2E) is V2E.Local + assert common.local_dimension_of(shared) is V2E.Local + with pytest.raises(TypeError, match="not a connectivity declaration"): + common.local_dimension_of(NeighborConnectivity) + + +class TestOffsetProvider: + def test_class_keys_become_tags(self): + table = _table() + assert common.as_tag_keyed_offset_provider({V2E: table}) == {V2E.Local.tag: table} + + def test_tag_keys_pass_through(self): + provider = {V2E.Local.tag: _table()} + assert common.as_tag_keyed_offset_provider(provider) is provider + + def test_bare_name_is_rejected(self): + with pytest.raises(TypeError, match="keyed by 'NeighborConnectivity' declarations"): + common.as_tag_keyed_offset_provider({"V2E": _table()}) + + def test_non_string_key_is_rejected(self): + with pytest.raises(TypeError, match="keyed by 'NeighborConnectivity' declarations"): + common.as_tag_keyed_offset_provider({V2E: _table(), Vertex: _table()}) + + def test_binding_twice_is_rejected(self): + with pytest.raises(ValueError, match="twice"): + common.as_tag_keyed_offset_provider({V2E: _table(), V2E.Local.tag: _table()}) + + def test_check_accepts_matching_tables(self): + common.check_offset_provider({V2E.Local.tag: _table()}) + common.check_offset_provider({V2E: _table().__gt_type__()}) + + def test_check_rejects_mismatching_tables(self): + with pytest.raises(ValueError, match="does not match its declaration"): + common.check_offset_provider({V2E.Local.tag: _table(codomain=Vertex)}) + + def test_sharing_connectivity_is_keyed_and_checked_by_its_own_tag(self): + provider = common.as_tag_keyed_offset_provider({V2E: _table(), V2EShared: _table()}) + assert set(provider) == {V2E.Local.tag, V2EShared.tag} + common.check_offset_provider(provider) + with pytest.raises(ValueError, match="'V2EShared' does not match its declaration"): + common.check_offset_provider({V2EShared.tag: _table(codomain=Vertex)}) + + def test_sharing_connectivities_need_the_same_skip_positions(self): + owner = _table(data=((0, 1, 2, 3), (1, 2, 3, 0))) + consistent = _table(data=((3, 2, 1, 0), (0, 3, 2, 1))) + inconsistent = _table(data=((3, 2, 1, -1), (0, 3, 2, 1))) + common.check_offset_provider({V2E: owner, V2EShared: consistent}, deep=True) + with pytest.raises(ValueError, match="different neighbor structure"): + common.check_offset_provider({V2E: owner, V2EShared: inconsistent}, deep=True) + # the call path does not read the tables + common.check_offset_provider({V2E: owner, V2EShared: inconsistent}) + + def test_the_connectivity_s_own_tag_is_rejected(self): + with pytest.raises(ValueError, match="whose key is the declaration itself"): + common.check_offset_provider({V2E.tag: _table()}) + + def test_a_checked_provider_is_remembered(self): + provider = {V2E: _table(codomain=Vertex)} + with pytest.raises(ValueError, match="does not match its declaration"): + common.check_offset_provider(provider) + # and a provider that passes is not checked twice + good = {V2E: _table()} + common.check_offset_provider(good) + common.check_offset_provider(good) + + def test_check_skips_undeclared_tags(self): + common.check_offset_provider({"some.hand.written.tag": _table()}) + + def test_missing_connectivity_error_is_actionable(self): + with pytest.raises(KeyError, match="keyed by 'NeighborConnectivity' declarations"): + common.get_offset({}, V2E.Local.tag) + + +class TestConnectivityKeyOver: + def _type(self, connectivity): + return _table_type(domain=(connectivity.origin, common.local_dimension_of(connectivity))) + + def test_owner_is_preferred(self): + ns = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + ) + shared = ns["V2EShared"] + provider = {shared.offset_tag: self._type(shared), V2E.offset_tag: self._type(V2E)} + assert common.connectivity_key_over(provider, V2E.Local) == V2E.offset_tag + assert common.connectivity_key_over(provider, V2E.Local.tag) == V2E.offset_tag + + def test_sharers_are_picked_independently_of_order(self): + ns = _declare( + """ + class SharedA(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + + class SharedB(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + ) + a, b = ns["SharedA"], ns["SharedB"] + forward = {a.offset_tag: self._type(a), b.offset_tag: self._type(b)} + backward = dict(reversed(forward.items())) + assert ( + common.connectivity_key_over(forward, V2E.Local) + == common.connectivity_key_over(backward, V2E.Local) + == min(a.offset_tag, b.offset_tag) + ) + + def test_nothing_bound(self): + with pytest.raises(KeyError, match="No connectivity over the local dimension"): + common.connectivity_key_over({E2V.offset_tag: self._type(E2V)}, V2E.Local) + + +def test_redefined_declaration_resolves_to_the_new_class(monkeypatch): + """Re-running a notebook cell redefines declarations under the same names.""" + import sys + import types as pytypes + + module = pytypes.ModuleType("_redefined_connectivity_module") + monkeypatch.setitem(sys.modules, module.__name__, module) + source = textwrap.dedent( + """ + from gt4py.next.common import DimensionIndex, LocalDimensionIndex, NeighborConnectivity + + class V(DimensionIndex): ... + class E(DimensionIndex): ... + class V2E(NeighborConnectivity[V, E], max_neighbors={n}): + class Local(LocalDimensionIndex): ... + """ + ) + exec(source.format(n=4), module.__dict__) + old = module.V2E + common.check_offset_provider({old: _table(domain=(module.V, old.Local), codomain=module.E)}) + + exec(source.format(n=2), module.__dict__) + new = module.V2E + assert common.resolve(new.Local.tag) is new.Local + common.check_offset_provider( + {new: _table(domain=(module.V, new.Local), codomain=module.E, data=((0, 1), (1, 0)))} + ) + with pytest.raises(ValueError, match="was the declaration redefined"): + common.check_neighbor_table( + new, _table(domain=(old.origin, old.Local), codomain=module.E, data=((0, 1), (1, 0))) + ) + + +def test_the_const_list_dimension_cannot_be_adopted(): + with pytest.raises(TypeError, match="cannot adopt"): + _declare( + """ + class C(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = ConstList + """ + ) diff --git a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py index 6ee13a358c..07f36a5737 100644 --- a/tests/next_tests/unit_tests/type_system_tests/test_type_info.py +++ b/tests/next_tests/unit_tests/type_system_tests/test_type_info.py @@ -14,6 +14,7 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, ) from gt4py.next.type_system import type_info, type_specifications as ts from gt4py.next.ffront import type_specifications as ts_ffront @@ -29,10 +30,10 @@ class JDim(DimensionIndex): ... class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2EDim(LocalDimensionIndex): ... class TDim(DimensionIndex): ... diff --git a/typing_tests/pyright_probes.py b/typing_tests/pyright_probes.py new file mode 100644 index 0000000000..e9be2423d9 --- /dev/null +++ b/typing_tests/pyright_probes.py @@ -0,0 +1,111 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +""" +Client code that has to type-check under *pyright*, checked by `nox -s test_typing_exports`. + +The cases in `test_next.yaml` run under mypy only, and the two checkers disagree about what +counts as a type: an annotated `Local` on a connectivity or its metaclass makes every +declaration's local dimension a *variable* for pyright, so `Field[Dims[V, V2E.Local], float]` +is rejected there while mypy accepts it (see ADR 0029). Everything here must be error-free. +""" + +from __future__ import annotations + +import typing + +from gt4py import next as gtx + + +class Vertex(gtx.DimensionIndex): ... + + +class Edge(gtx.DimensionIndex): ... + + +class Cell(gtx.DimensionIndex): ... + + +class CellEdge(gtx.DimensionIndex): ... + + +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): + class Local(gtx.LocalDimensionIndex): ... + + +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + +#: A flattened sparse pattern sharing `C2E`'s neighbor axis. +class C2CE(gtx.NeighborConnectivity[Cell, CellEdge]): + Local: typing.TypeAlias = C2E.Local + + +class LsqCoeff(gtx.LocalDimensionIndex, size=3): ... + + +#: A declaration adopting a local dimension declared at module level. +class V2EAdopted(gtx.NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = LsqCoeff + + +def nested_local(sparse: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + +def shared_local(sparse: gtx.Field[gtx.Dims[Cell, C2CE.Local], gtx.float64]) -> None: ... + + +def adopted_local(sparse: gtx.Field[gtx.Dims[Vertex, V2EAdopted.Local], gtx.float64]) -> None: ... + + +def a_shared_local_is_its_owners( + owned: gtx.Field[gtx.Dims[Cell, C2E.Local], gtx.float64], + shared: gtx.Field[gtx.Dims[Cell, C2CE.Local], gtx.float64], +) -> None: + shared_local(owned) # the two spellings are one type + shared_local(shared) + + +def an_adopted_local_is_the_adopted_one( + coefficients: gtx.Field[gtx.Dims[Vertex, LsqCoeff], gtx.float64], +) -> None: + adopted_local(coefficients) + + +L = typing.TypeVar("L", bound=gtx.LocalDimensionIndex) + + +def local_of(connectivity: type[gtx.NeighborConnectivity]) -> type[gtx.LocalDimensionIndex]: + # generic code names a local dimension through the accessor, not through `conn.Local` + return gtx.local_dimension_of(connectivity) + + +def first_neighbor(sparse: gtx.Field[gtx.Dims[Vertex, L], gtx.float64]) -> type[L]: + raise NotImplementedError + + +def generic_local(sparse: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: + typing.assert_type(first_neighbor(sparse), type[V2E.Local]) + + +@gtx.field_operator +def reduce_over_a_local_dimension( + a: gtx.Field[gtx.Dims[Edge], gtx.float64], +) -> gtx.Field[gtx.Dims[Vertex], gtx.float64]: + return gtx.neighbor_sum(a(V2E), axis=V2E.Local) + + +@gtx.field_operator +def shift_by_a_dimension( + a: gtx.Field[gtx.Dims[KDim], gtx.float64], +) -> gtx.Field[gtx.Dims[KDim], gtx.float64]: + return a(KDim + 1) diff --git a/typing_tests/pyrightconfig.json b/typing_tests/pyrightconfig.json new file mode 100644 index 0000000000..4f44d9cfb8 --- /dev/null +++ b/typing_tests/pyrightconfig.json @@ -0,0 +1,5 @@ +{ + "typeCheckingMode": "standard", + "reportMissingImports": "error", + "reportMissingTypeStubs": "none" +} diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 20ed020333..e3f47b0f1e 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -275,3 +275,74 @@ main: | import xarray a: xarray.NamedArray + + - case: neighbor_connectivity_declaration + main: | + from __future__ import annotations + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6): + class Local(gtx.LocalDimensionIndex): ... + + def sparse(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + reveal_type(V2E.Local) + reveal_type(V2E.Local.owner) + reveal_type(V2E[1]) + out: | + main:12:13: note: Revealed type is "def (value: int) -> main.V2E.Local" + main:13:13: note: Revealed type is "type[gt4py.next.common.NeighborConnectivity[Any, Any]] | None" + main:14:13: note: Revealed type is "gt4py.next.common.Connectivity[Any, Any]" + + - case: neighbor_connectivity_locals_are_distinct + main: | + from __future__ import annotations + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + class Cell(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + class V2C(gtx.NeighborConnectivity[Vertex, Cell]): + class Local(gtx.LocalDimensionIndex): ... + + def takes_v2e(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: ... + + def caller(a: gtx.Field[gtx.Dims[Vertex, V2C.Local], gtx.float64]) -> None: + takes_v2e(a) + out: | + main:17:15: error: Argument 1 to "takes_v2e" has incompatible type "Field[Dims[Vertex, main.V2C.Local], float]"; expected "Field[Dims[Vertex, main.V2E.Local], float]" [arg-type] + main:17:15: note: "Field[Dims[Vertex, Local], float].__call__" has type "def __call__(self, index_field: Connectivity[Any, Any] | type[NeighborConnectivity[Any, Any]], *args: Connectivity[Any, Any] | type[NeighborConnectivity[Any, Any]]) -> Field[Any, Any]" + + - case: neighbor_connectivity_generic_local + main: | + from __future__ import annotations + import typing + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + L = typing.TypeVar("L", bound=gtx.LocalDimensionIndex) + + def local_of(conn: type[gtx.NeighborConnectivity]) -> type[gtx.LocalDimensionIndex]: + # `conn.Local` is not annotated, so that a declaration's `Local` stays a type; generic + # code reads it through the accessor + return gtx.local_dimension_of(conn) + + def first(a: gtx.Field[gtx.Dims[Vertex, L], gtx.float64]) -> type[L]: + raise NotImplementedError + + def caller(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: + reveal_type(first(a)) + out: | + main:22:17: note: Revealed type is "type[main.V2E.Local]" diff --git a/uv.lock b/uv.lock index d024752ae3..d3c356efda 100644 --- a/uv.lock +++ b/uv.lock @@ -1441,6 +1441,7 @@ typing = [ ] typing-exports = [ { name = "mypy", extra = ["faster-cache"] }, + { name = "pyright" }, { name = "pytest-mypy-plugins" }, { name = "types-decorator" }, { name = "types-docutils" }, @@ -1596,6 +1597,7 @@ typing = [ ] typing-exports = [ { name = "mypy", extras = ["faster-cache"], specifier = ">=1.13.0" }, + { name = "pyright", specifier = ">=1.1.400" }, { name = "pytest-mypy-plugins", specifier = ">=4.0.0" }, { name = "types-decorator", specifier = ">=5.1.8" }, { name = "types-docutils", specifier = ">=0.21.0" }, @@ -3124,6 +3126,19 @@ version = "2.1" source = { registry = "https://pypi.org/simple" } sdist = { url = "https://files.pythonhosted.org/packages/bc/7c/d724ef1ec3ab2125f38a1d53285745445ec4a8f19b9bb0761b4064316679/pyreadline-2.1.zip", hash = "sha256:4530592fc2e85b25b1a9f79664433da09237c1a270e4d78ea5aa3a2c7229e2d1", size = 109189, upload-time = "2015-09-16T08:24:48.745Z" } +[[package]] +name = "pyright" +version = "1.1.414" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nodeenv" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/e1/1b/244c7b710031ada80f27e579ec20d28a2285dfc318fed0339866b1047f12/pyright-1.1.414.tar.gz", hash = "sha256:523c0a97c60da6333234955c277730c9cf4f5bd6d5399e7b7d2b0fc5d3599524", size = 4154638, upload-time = "2026-09-10T12:26:53.181Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d7/ba/18b6e682ead424ad24bcc134339ae5d1b931cd9ae260540592a058a91279/pyright-1.1.414-py3-none-any.whl", hash = "sha256:2a6b4b3298c9eec174c5ed83bd338de6eee82df2992f3e1930e6199d381be36f", size = 6225049, upload-time = "2026-09-10T12:26:51.427Z" }, +] + [[package]] name = "pytest" version = "9.1.1"