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..150976307b --- /dev/null +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -0,0 +1,158 @@ +--- +tags: [] +--- + +# Connectivities as Types + +- **Status**: proposed +- **Authors**: Enrique González Paredes (@egparedes) +- **Created**: 2026-09-21 +- **Updated**: 2026-09-21 + +A neighbor connectivity is declared as a **class**, and its local dimension as a +class **nested** in it: + +```python +class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5): + class Local(gtx.LocalDimensionIndex): ... + + +@gtx.field_operator +def f(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return neighbor_sum(a(V2E), axis=V2E.Local) + a(V2E[0]) +``` + +The declaration is written in DSL code, owns its local dimension, and states the +constraints 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. The check is explicit for as long as + offset providers are keyed by tag strings, since nothing then connects a + provider entry to a declaration; it becomes automatic with class-keyed + providers. + +### `NeighborConnectivity` is not a `Connectivity` + +`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 still find + the owner's table by the local dimension's tag, for its neighbor structure, so + the owner has to be bound too. + `V2E.Local` inside DSL code types as that local dimension, and + `FieldOffset.Local` names the same thing on a legacy offset, so the spelling + works for both. The other frontend touch points treat the class like the + `FieldOffset` it derives: grid-type deduction (`transform_utils`, `past_to_itir`) + counts it as unstructured, and embedded `premap` accepts it. `V2E[i]` subscripts the + metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. + +## Consequences + +- 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` remains during migration; a `FieldOffset` and a + `NeighborConnectivity` sharing a local dimension are interchangeable. + +## Alternatives considered + +- **The local dimension generated by the metaclass**, e.g. `V2E.Local` created + from `V2E`'s name. Type checkers cannot see a generated class, so it could not + be used in `Field[Dims[Vertex, V2E.Local], ...]`. +- **Neighbor counts as type parameters.** Python has no integer type parameters, + and a `Literal[6]` argument would add a type parameter nothing statically uses. +- **`NeighborConnectivity` as a `Connectivity` subclass.** Mixes the + declaration with the data protocol; see above. diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index 9c93fd9aa0..e663e462f3 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -23,6 +23,7 @@ Writing a new ADR is simple: - [0023 - Fingerprinting](0023-Fingerprinting.md) - [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md) - [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md) +- [0029 - Connectivities as Types](0029-Connectivities_As_Types.md) ### Frontend and Parsing #frontend diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 1e1cbfc280..08c37acb03 100644 --- a/docs/user/next/QuickstartGuide.md +++ b/docs/user/next/QuickstartGuide.md @@ -218,28 +218,28 @@ edge_values = gtx.as_field([EdgeDim], np.zeros((12,))) +++ -You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _field offset_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements. +You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _connectivity_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements. To understand this transform, you can look at the edge-to-cell connectivity table `edge_to_cell_table` listed above. This table has the same shape as the output of the transform, that is, one dimension over the edges and another _local_ dimension. The table stores indices into a field over cells, the transform essentially gives you another field where the indices have been replaced with the values in the cell field at the corresponding indices. Another way to look at it is that transform uses the edge-to-cell connectivity to look up all the cell neighbors of edges, and associates the values of those neighbor cells with each edge. -You can use the field offset `E2C` below to transform a field over cells to a field over edges using the edge-to-cell connectivities: +You can use the connectivity `E2C` declared below to transform a field over cells to a field over edges using the edge-to-cell connectivities. It is declared as a class: for each edge (`EdgeDim`), a list of neighbor cells (`CellDim`), indexed by its nested local dimension `E2C.Local`: ```{code-cell} ipython3 -class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... -E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim)) +class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + class Local(gtx.LocalDimensionIndex): ... ``` -The field offset is named by its local dimension's `tag`, and the offset provider below is keyed by the same `tag`, so all three refer to one connectivity. Note that the field offset does not contain the actual connectivity table, that's provided through an _offset provider_: +Note that the declaration does not contain the actual connectivity table, that's provided through an _offset provider_, keyed by the local dimension's `tag`: ```{code-cell} ipython3 -E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2CDim], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) +E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2C.Local], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) ``` The field operator `nearest_cell_to_edge` below shows an example of applying this transform. There is a little twist though: the subscript in `E2C[0]` means that only the value of the first connected cell is taken, the second (if exists) is ignored. -Pay attention to the syntax where the field offset `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. +Pay attention to the syntax where the connectivity `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. ```{code-cell} ipython3 @gtx.field_operator @@ -250,7 +250,7 @@ def nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64]) -> gtx. def run_nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): nearest_cell_to_edge(cell_values, out=out) -run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider}) print("0th adjacent cell's value: {}".format(edge_values.asnumpy())) ``` @@ -265,19 +265,19 @@ Running the above snippet results in the following edge field: #### Using reductions on connected mesh elements -Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2CDim) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2CDim` dimension. +Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2C.Local) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2C.Local` dimension. ```{code-cell} ipython3 @gtx.field_operator def sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64]) -> gtx.Field[Dims[EdgeDim], float64]: - # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2CDim], float64] - return neighbor_sum(cells(E2C), axis=E2CDim) + # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2C.Local], float64] + return neighbor_sum(cells(E2C), axis=E2C.Local) @gtx.program def run_sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): sum_adjacent_cells(cells, out=out) -run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider}) print("sum of adjacent cells: {}".format(edge_values.asnumpy())) ``` @@ -376,13 +376,13 @@ print("where nested tuple return: {}".format(((result_1.asnumpy(), result_2.asnu #### Implementing the pseudo-laplacian -As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to define the local dimension, the field offset, and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: +As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to declare the connectivity and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: ```{code-cell} ipython3 -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=EdgeDim, target=(CellDim, C2EDim)) +class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + class Local(gtx.LocalDimensionIndex): ... -C2E_offset_provider = gtx.as_connectivity([CellDim, C2EDim], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) +C2E_offset_provider = gtx.as_connectivity([CellDim, C2E.Local], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) ``` **Weights of edge differences:** @@ -410,7 +410,7 @@ edge_weights = np.array([ [0, -1, -1], # cell 5 ], dtype=np.float64) -edge_weight_field = gtx.as_field([CellDim, C2EDim], edge_weights) +edge_weight_field = gtx.as_field([CellDim, C2E.Local], edge_weights) ``` Now you have everything to implement the pseudo-laplacian. Its field operator requires the cell field and the edge weights as inputs, and outputs a cell field of the same shape as the input. @@ -422,9 +422,9 @@ The second lines first creates a temporary field using `edge_differences(C2E)`, ```{code-cell} ipython3 @gtx.field_operator def pseudo_lap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: edges = cells(E2C[0]) # type: gtx.Field[Dims[EdgeDim], float64] - return neighbor_sum(edges(C2E) * edge_weights, axis=C2EDim) + return neighbor_sum(edges(C2E) * edge_weights, axis=C2E.Local) ``` The program itself is just a shallow wrapper over the `pseudo_lap` field operator. The significant part is how offset providers for both the edge-to-cell and cell-to-edge connectivities are supplied when the program is called: @@ -432,7 +432,7 @@ The program itself is just a shallow wrapper over the `pseudo_lap` field operato ```{code-cell} ipython3 @gtx.program def run_pseudo_laplacian(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64], + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64], out : gtx.Field[Dims[CellDim], float64]): pseudo_lap(cells, edge_weights, out=out) @@ -441,7 +441,7 @@ result_pseudo_lap = gtx.as_field([CellDim], np.zeros(shape=(6,))) run_pseudo_laplacian(cell_values, edge_weight_field, result_pseudo_lap, - offset_provider={E2CDim.tag: E2C_offset_provider, C2EDim.tag: C2E_offset_provider}) + offset_provider={E2C.Local.tag: E2C_offset_provider, C2E.Local.tag: C2E_offset_provider}) print("pseudo-laplacian: {}".format(result_pseudo_lap.asnumpy())) ``` @@ -451,7 +451,7 @@ As a closure, here is an example of chaining field operators, which is very simp ```{code-cell} ipython3 @gtx.field_operator def pseudo_laplap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: return pseudo_lap(pseudo_lap(cells, edge_weights), edge_weights) ``` diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb index b0a1980d0f..21bf2d25d8 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb @@ -126,7 +126,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb index 573ee6a44e..86c8d33ac7 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb @@ -131,7 +131,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb index 2b422b1823..fb2282ab22 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb @@ -123,7 +123,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb index 85044b989f..43196507ce 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb @@ -136,7 +136,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E.Local.tag: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb index dc321f1bdd..b99c6f6d3f 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb @@ -147,7 +147,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E.Local.tag: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb index 251fe8239a..de040ccb93 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb @@ -152,7 +152,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E.Local.tag: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb index 30f568de6f..174699e350 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb @@ -293,10 +293,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E.Local.tag: c2e_connectivity,\n", + " V2E.Local.tag: v2e_connectivity,\n", + " E2V.Local.tag: e2v_connectivity,\n", + " E2C.Local.tag: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb index eaeb8c7b02..81836edbd5 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb @@ -314,10 +314,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E.Local.tag: c2e_connectivity,\n", + " V2E.Local.tag: v2e_connectivity,\n", + " E2V.Local.tag: e2v_connectivity,\n", + " E2C.Local.tag: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb index b278cee26d..edd65ac9e5 100644 --- a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb @@ -169,7 +169,7 @@ " kappa,\n", " dt,\n", " out=(divergence_gt4py_1, divergence_gt4py_2),\n", - " offset_provider={E2C2V.value: e2c2v_connectivity, V2E.value: v2e_connectivity},\n", + " offset_provider={E2C2V.Local.tag: e2c2v_connectivity, V2E.Local.tag: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py_1.asnumpy(), divergence_ref_1)\n", diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index c398524538..07fb984328 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -11,7 +11,13 @@ import gt4py.next as gtx from gt4py.next.iterator.embedded import MutableLocatedField from gt4py.next import neighbor_sum, where, Dims -from gt4py.next import Dimension, DimensionIndex, DimensionKind, FieldOffset +from gt4py.next import ( + Dimension, + DimensionIndex, + DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, +) from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( run_gtfn as gtfn_cpu, @@ -389,31 +395,36 @@ class E(DimensionIndex): ... class K(DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... -class C2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E(NeighborConnectivity[C, E]): + class Local(LocalDimensionIndex): ... -C2E = FieldOffset(C2EDim.tag, source=E, target=(C, C2EDim)) +C2EDim = C2E.Local -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(NeighborConnectivity[V, E]): + class Local(LocalDimensionIndex): ... -V2E = FieldOffset(V2EDim.tag, source=E, target=(V, V2EDim)) +V2EDim = V2E.Local -class E2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(NeighborConnectivity[E, V]): + class Local(LocalDimensionIndex): ... -E2V = FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) +E2VDim = E2V.Local -class E2CDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C(NeighborConnectivity[E, C]): + class Local(LocalDimensionIndex): ... -E2C = FieldOffset(E2CDim.tag, source=C, target=(E, E2CDim)) +E2CDim = E2C.Local -class E2C2VDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2V(NeighborConnectivity[E, V]): + class Local(LocalDimensionIndex): ... -E2C2V = FieldOffset(E2C2VDim.tag, source=V, target=(E, E2C2VDim)) +E2C2VDim = E2C2V.Local diff --git a/docs/user/next/workshop/slides/slides_2.ipynb b/docs/user/next/workshop/slides/slides_2.ipynb index db8f370abc..f16f1560bc 100644 --- a/docs/user/next/workshop/slides/slides_2.ipynb +++ b/docs/user/next/workshop/slides/slides_2.ipynb @@ -273,10 +273,11 @@ "metadata": {}, "outputs": [], "source": [ - "class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...\n", + "class E2C(gtx.NeighborConnectivity[Edge, Cell]):\n", + " class Local(gtx.LocalDimensionIndex): ...\n", "\n", "\n", - "E2C = gtx.FieldOffset(E2CDim.tag, source=Cell, target=(Edge, E2CDim))" + "E2CDim = E2C.Local" ] }, { diff --git a/noxfile.py b/noxfile.py index c8091de09e..31327f076d 100755 --- a/noxfile.py +++ b/noxfile.py @@ -349,6 +349,9 @@ def test_typing_exports(session: nox.Session) -> None: "typing_tests", *session.posargs, ) + # A second checker, on code that must type-check for a downstream user: mypy and pyright + # disagree about what counts as a type, which the mypy-only cases above cannot catch. + session.run("pyright", "--project", "typing_tests", "typing_tests/pyright_probes.py") # -- DaCe codegen determinism check -- diff --git a/pyproject.toml b/pyproject.toml index 64f52fdb8b..89a743ce18 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -64,6 +64,7 @@ typing = [ typing_exports = [ # to test typing with gt4py in downstream code {include-group = "typing"}, + 'pyright>=1.1.400', # the second checker: it disagrees with mypy about what counts as a type 'types-six>=1.17.0.20251009', # can not let mypy auto-install types as that leads to unexpected stderr output (which means test failure) 'pytest-mypy-plugins>=4.0.0', # pytest plugin for running mypy on code snippets "xarray>=2024.1.0" # one of the regression tests requires xarray diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index e665024d7d..94690b7d0e 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, ) @@ -120,6 +123,8 @@ "Dimension", "DimensionIndex", "DimensionKind", + "LocalDimensionIndex", + "NeighborConnectivity", "Staggered", "resolve", "Dims", @@ -132,6 +137,7 @@ "unit_range", "UnitRange", "is_staggered", + "local_dimension_of", "flip_staggered", "as_non_staggered", # from constructors diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 89014c25c0..d91bbff004 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 @@ -1027,7 +1036,9 @@ def asnumpy(self) -> np.ndarray: ... def as_scalar(self) -> core_defs.ScalarT: ... @abc.abstractmethod - def premap(self, index_field: Connectivity | fbuiltins.FieldOffset) -> Field: ... + def premap( + self, index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity] + ) -> Field: ... @abc.abstractmethod def restrict(self, item: AnyIndexSpec) -> Self: ... @@ -1035,8 +1046,8 @@ def restrict(self, item: AnyIndexSpec) -> Self: ... @abc.abstractmethod def __call__( self, - index_field: Connectivity | fbuiltins.FieldOffset, - *args: Connectivity | fbuiltins.FieldOffset, + index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], ) -> Field: ... @abc.abstractmethod @@ -1219,7 +1230,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 @@ -1557,8 +1571,8 @@ def inverse_image(self, image_range: UnitRange | NamedRange) -> Sequence[NamedRa def premap( self, - index_field: Connectivity | fbuiltins.FieldOffset, - *args: Connectivity | fbuiltins.FieldOffset, + index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], ) -> Connectivity: raise NotImplementedError() @@ -1627,8 +1641,8 @@ def promote_dims(*dims_list: Sequence[Dimension]) -> list[Dimension]: >>> class I(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... >>> class J(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... >>> class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... - >>> class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... - >>> class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... + >>> class E2V(LocalDimensionIndex): ... + >>> class E2C(LocalDimensionIndex): ... >>> promote_dims([J, K], [I, K]) == [I, J, K] True >>> promote_dims([K, J], [I, K]) @@ -1731,6 +1745,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 +1807,6 @@ def __init_subclass__(cls, /, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) -class ConstListDim(DimensionIndex, kind=DimensionKind.LOCAL): - """ - The local dimension of a list whose length is known at compile time (`make_const_list`). - - Declared here, once, because it must be a *single* class. It used to be built - independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless - while dimensions compared by `(name, kind)` -- the two instances were equal. Under - nominal identity (ADR 0028) two declarations would be two different dimensions, and the - `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s - built by embedded execution. - - TODO: becomes an owner-less local dimension with an explicit size, generalising this from - length 1 to length *n*, once local dimensions know their connectivity. - """ - - __slots__ = () - - def _reduce_staggered(cls: StaggeredMeta) -> Any: """ Pickle a staggered dimension through its base, falling back to by-reference. @@ -1878,3 +1876,369 @@ 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 ConstListDim(LocalDimensionIndex): + """ + The local dimension of a list whose length is known at compile time (`make_const_list`). + + Declared here, once, because it must be a *single* class. It used to be built + independently in `iterator/embedded.py` and in the DaCe lowering, which was harmless + while dimensions compared by `(name, kind)` -- the two instances were equal. Under + nominal identity (ADR 0028) two declarations would be two different dimensions, and the + `offset_type == _CONST_DIM` checks in the DaCe lowering would stop matching `ListType`s + built by embedded execution. + + TODO: becomes an owner-less local dimension with an explicit size, generalising this from + length 1 to length *n*, once local dimensions know their connectivity. + """ + + __slots__ = () + + +class ConnectivityMeta(type): + """ + Metaclass of `NeighborConnectivity` declarations. + + 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.__gt_field_offset__()[int(item)] + if "Local" in cls.__dict__: + raise TypeError( + f"'{cls.__qualname__}[{item!r}]': a connectivity is indexed by an integer" + " neighbor position." + ) + # A metaclass `__getitem__` shadows `__class_getitem__`, so type-parameter + # subscription (`NeighborConnectivity[V, E]`) has to be forwarded explicitly. + return cast(Any, cls).__class_getitem__(item) + + def __repr__(cls) -> str: + return cls.tag + + def __str__(cls) -> str: + return cls.__qualname__ + + def __gt_type__(cls) -> Any: + return cls.__gt_field_offset__().__gt_type__() + + def __gt_field_offset__(cls) -> Any: + """ + The `FieldOffset` equivalent to this connectivity, tagged with `offset_tag`. + """ + from gt4py.next.ffront import fbuiltins + + if (field_offset := cls.__dict__.get("_field_offset")) is None: + field_offset = fbuiltins.FieldOffset( + cls.offset_tag, + source=cls.codomain, + target=(cls.origin, cls._local()), + _derived=True, + ) + type.__setattr__(cls, "_field_offset", field_offset) + return field_offset + + +class NeighborConnectivity[Origin: DimensionIndex, Codomain: DimensionIndex]( + metaclass=ConnectivityMeta +): + """ + Declare a neighbor connectivity: for each `Origin` element, a list of `Codomain` neighbors. + + The declaration names the connectivity's local dimension -- its nested `Local` class -- + and optionally its neighbor counts. It holds no data: the neighbor table is bound at call + time through the offset provider. `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 ConstListDim: + raise TypeError( + f"'{name}' cannot adopt '{ConstListDim.__qualname__}': it is the local dimension" + " of 'make_const_list' results and belongs to no connectivity." + ) + max_neighbors = _check_neighbor_count(cls, "max_neighbors", max_neighbors) + min_neighbors = _check_neighbor_count(cls, "min_neighbors", min_neighbors) + # NOTE: a declaration whose tag is the owner's is a *redefinition* of it (a re-run + # notebook cell), not a second connectivity sharing the local dimension, so it takes + # ownership over again. Ownership of a local dimension that is adopted, rather than + # nested, otherwise goes to whoever declares first. + if local.owner is not None and local.owner.tag != cls.tag: + # Sharing another connectivity's local dimension: the neighbor structure is the + # owner's, including its counts, and the sharing connectivity is named by its own tag. + if (max_neighbors, min_neighbors) != (None, None) and ( + max_neighbors, + min_neighbors, + ) != (local.max_neighbors, local.min_neighbors): + raise TypeError( + f"'{name}' shares the local dimension of '{local.owner.__qualname__}', whose" + " neighbor counts are declared by its owner." + ) + cls.origin, cls.codomain = origin, codomain + return + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("min_neighbors", min_neighbors), + ): + # 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}'") + expected_domain = (connectivity.origin, local) + if tuple(table_type.domain) != expected_domain: + fail( + f"its domain is '({', '.join(map(str, table_type.domain))})'," + f" expected '({', '.join(map(str, expected_domain))})'" + ) + if table_type.codomain is not connectivity.codomain: + fail(f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'") + if 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}" + ) 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/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index 0e3aaeab8b..f8c422f134 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -239,7 +239,9 @@ def as_scalar(self) -> core_defs.ScalarT: def premap( self: NdArrayField, - *connectivities: common.Connectivity | fbuiltins.FieldOffset, + *connectivities: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], ) -> NdArrayField: """ Rearrange the field content using the provided connectivities (index mappings). @@ -314,7 +316,10 @@ def premap( codomains_counter: collections.Counter[common.Dimension] = collections.Counter() for connectivity in connectivities: - # For neighbor reductions, a FieldOffset is passed instead of an actual Connectivity + # For neighbor reductions, a FieldOffset or a connectivity declaration is passed + # instead of an actual Connectivity + if isinstance(connectivity, common.ConnectivityMeta): + connectivity = connectivity.__gt_field_offset__() if not isinstance(connectivity, common.Connectivity): assert isinstance(connectivity, fbuiltins.FieldOffset) connectivity = connectivity.as_connectivity_field() @@ -366,8 +371,10 @@ def premap( def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: return functools.reduce( lambda field, current_index_field: field.premap(current_index_field), diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 7474cd1406..0916417a5f 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -11,6 +11,7 @@ import inspect import math import operator +import warnings from builtins import bool, float, int, tuple # noqa: A004 shadowing a Python built-in from types import UnionType from typing import ( @@ -484,18 +485,39 @@ class FieldOffset(runtime.Offset): value: str source: common.Dimension target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] + #: Set when derived from a `NeighborConnectivity` declaration, which is not deprecated. + _derived: bool = dataclasses.field(default=False, repr=False, compare=False, kw_only=True) @functools.cached_property def _cache(self) -> dict: return {} def __post_init__(self) -> None: - if len(self.target) == 2 and self.target[1].kind != common.DimensionKind.LOCAL: - raise ValueError("Second dimension in offset must be a local dimension.") + if len(self.target) == 2: + if self.target[1].kind != common.DimensionKind.LOCAL: + raise ValueError("Second dimension in offset must be a local dimension.") + if not self._derived: + warnings.warn( + "Declaring an unstructured connectivity with 'FieldOffset' is deprecated;" + " declare a 'NeighborConnectivity' class instead (see ADR 0029):\n" + " class V2E(gtx.NeighborConnectivity[Vertex, Edge]):\n" + " class Local(gtx.LocalDimensionIndex): ...", + DeprecationWarning, + stacklevel=3, + ) def __gt_type__(self) -> ts.OffsetType: return ts.OffsetType(source=self.source, target=self.target, tag=self.value) + @property + def Local(self) -> common.Dimension: + """The local dimension, as `V2E.Local` names it on a `NeighborConnectivity`.""" + if len(self.target) != 2: + raise AttributeError( + f"'{self.value}' is a Cartesian offset and has no local dimension." + ) + return self.target[1] + def __getitem__(self, offset: int) -> common.Connectivity: """Serve as a connectivity factory.""" from gt4py.next import embedded # avoid circular import diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 3d09d7ecb4..85eaafe610 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: diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 1b6cefc6b2..131647baf4 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -74,7 +74,7 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> stages.CompilableProgramDef: """ all_closure_vars = transform_utils._get_closure_vars_recursively(inp.data.closure_vars) offsets_and_dimensions = transform_utils._filter_closure_vars_by_type( - all_closure_vars, fbuiltins.FieldOffset, common.DimensionMeta + all_closure_vars, fbuiltins.FieldOffset, common.ConnectivityMeta, common.DimensionMeta ) grid_type = transform_utils._deduce_grid_type( inp.data.grid_type, offsets_and_dimensions.values() diff --git a/src/gt4py/next/ffront/transform_utils.py b/src/gt4py/next/ffront/transform_utils.py index 09c9d4b9ee..4e24d83881 100644 --- a/src/gt4py/next/ffront/transform_utils.py +++ b/src/gt4py/next/ffront/transform_utils.py @@ -47,7 +47,9 @@ def _filter_closure_vars_by_type(closure_vars: dict[str, Any], *types: type) -> def _deduce_grid_type( requested_grid_type: Optional[common.GridType], - offsets_and_dimensions: Iterable[fbuiltins.FieldOffset | common.Dimension], + offsets_and_dimensions: Iterable[ + fbuiltins.FieldOffset | type[common.NeighborConnectivity] | common.Dimension + ], ) -> common.GridType: """ Derive grid type from actually occurring dimensions and check against optional user request. @@ -59,7 +61,9 @@ def _deduce_grid_type( deduced_grid_type = common.GridType.CARTESIAN for o in offsets_and_dimensions: - if isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o): + if isinstance(o, common.ConnectivityMeta) or ( + isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o) + ): deduced_grid_type = common.GridType.UNSTRUCTURED break if isinstance(o, common.DimensionMeta) and o.kind == common.DimensionKind.LOCAL: diff --git a/src/gt4py/next/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..5a9e4011cc 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -154,7 +154,12 @@ def ndarray(self) -> core_defs.NDArrayObject: def asnumpy(self) -> np.ndarray: raise NotImplementedError - def premap(self, index_field: common.Connectivity | fbuiltins.FieldOffset) -> common.Field: + def premap( + self, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + ) -> common.Field: raise NotImplementedError def restrict( # type: ignore[override] @@ -171,8 +176,10 @@ def as_scalar(self) -> xtyping.Never: def __call__( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError() @@ -1152,8 +1159,10 @@ def as_scalar(self) -> core_defs.IntegralScalar: def premap( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1293,8 +1302,10 @@ def asnumpy(self) -> np.ndarray: def premap( self, - index_field: common.Connectivity | fbuiltins.FieldOffset, - *args: common.Connectivity | fbuiltins.FieldOffset, + index_field: common.Connectivity + | fbuiltins.FieldOffset + | type[common.NeighborConnectivity], + *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1389,10 +1400,10 @@ def constant_field(value: Any, dtype_like: Optional[core_defs.DTypeLike] = None) @builtins.shift.register(EMBEDDED) def shift( - *offsets: Union[runtime.Offset, OffsetPart], + *offsets: Union[runtime.Offset, type[common.NeighborConnectivity], OffsetPart], ) -> Callable[[ItIterator], ItIterator]: def impl(it: ItIterator) -> ItIterator: - return it.shift(*list(o.value if isinstance(o, runtime.Offset) else o for o in offsets)) + return it.shift(*list(_as_offset_tag(o) for o in offsets)) return impl @@ -1437,9 +1448,20 @@ def __gt_type__(self) -> ts.ListType: ) +def _as_offset_tag( + offset: runtime.Offset | type[common.NeighborConnectivity] | OffsetPart, +) -> OffsetPart: + if isinstance(offset, common.ConnectivityMeta): + return common.local_dimension_of(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 = ( + offset.__gt_field_offset__() if isinstance(offset, common.ConnectivityMeta) else offset + ) + offset_str = _as_offset_tag(field_offset) assert isinstance(offset_str, str) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None @@ -1451,7 +1473,7 @@ def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: for i in range(connectivity.__gt_type__().max_neighbors) if (shifted := it.shift(offset_str, i)).can_deref() ), - offset=offset, + offset=field_offset, ) diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 4e5b8d4b5a..094c85925f 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -153,6 +153,8 @@ def make_node(o): # it, see `execute_shift`); decide whether to fold it into the shift value or forbid it. assert o.offset == 0 return im.cartesian_offset(o.domain_dim, o.codomain) + if isinstance(o, common.ConnectivityMeta): + return OffsetLiteral(value=common.local_dimension_of(o).tag) if callable(o): if o.__name__ == "": return lambdadef(o) diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 7415bcb7a1..17b0949c6a 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -406,7 +406,7 @@ def _canonicalize_nb_fields( Examples: >>> class Vertex(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> input_field = ts.FieldType( ... dims=[ ... Vertex, diff --git a/src/gt4py/next/type_system/type_info.py b/src/gt4py/next/type_system/type_info.py index 2ce4672a77..da0cf84b1a 100644 --- a/src/gt4py/next/type_system/type_info.py +++ b/src/gt4py/next/type_system/type_info.py @@ -408,7 +408,7 @@ def is_local_field(type_: ts.FieldType) -> bool: Examples: >>> class V(common.DimensionIndex): ... - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> is_local_field( ... ts.FieldType(dims=[V, V2E], dtype=ts.ScalarType(kind=ts.ScalarKind.INT64)) ... ) @@ -586,7 +586,7 @@ def promote( >>> promoted.dims == [I, J, K] and promoted.dtype == dtype True - >>> class V2E(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... + >>> class V2E(common.LocalDimensionIndex): ... >>> list_dtype = ts.ListType(element_type=dtype, offset_type=V2E) >>> promote( ... ts.FieldType(dims=[I], dtype=list_dtype), diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 464bf7fa1f..2efc73b47c 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -35,7 +35,17 @@ # Both modules used to declare their own `Dimension("Vertex")` etc., which compared equal; under # nominal identity (ADR 0028) that would be two different dimensions, and tests that mix a # `toy_connectivity` connectivity with a `cases_utils` mesh would silently stop matching. -from next_tests.toy_connectivity import C2EDim, Cell, E2VDim, Edge, V2EDim, Vertex +from next_tests.toy_connectivity import ( + C2E, + C2EDim, + Cell, + E2V, + E2VDim, + Edge, + V2E, + V2EDim, + Vertex, +) __all__ = [ @@ -184,13 +194,11 @@ class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) -class C2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2V(gtx.NeighborConnectivity[Cell, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) -C2V = gtx.FieldOffset(C2VDim.tag, source=Vertex, target=(Cell, C2VDim)) +C2VDim = C2V.Local size = 10 @@ -308,28 +316,28 @@ def simple_mesh(allocator) -> MeshDescriptor: e2v_arr = np.asarray(e2v_arr, dtype=gtx.IndexType) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E.Local.tag: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 4}, codomain=Edge, data=v2e_arr, skip_value=None, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V.Local.tag: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V.Local.tag: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 4}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E.Local.tag: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 4}, codomain=Edge, data=c2e_arr, @@ -403,28 +411,28 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E.Local.tag: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 5}, codomain=Edge, data=v2e_arr, skip_value=common._DEFAULT_SKIP_VALUE, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V.Local.tag: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V.Local.tag: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 3}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E.Local.tag: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 3}, codomain=Edge, data=c2e_arr, diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py index fa65a38b40..f03f970507 100644 --- a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py +++ b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py @@ -42,10 +42,11 @@ class Cell(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2EDim = C2E.Local C2E_TABLE = np.array( [ diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py index 4a642f5cc4..57a2af8dfc 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py @@ -252,7 +252,7 @@ def test_compile_unstructured(unstructured_case, compile_testee_unstructured): compile_testee_unstructured(*args, offset_provider=unstructured_case.offset_provider, **kwargs) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E.Local.tag].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), @@ -317,7 +317,7 @@ def test_compile_unstructured_for_two_offset_providers( *args, offset_provider=unstructured_case.offset_provider, **kwargs ) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E.Local.tag].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), diff --git a/tests/next_tests/integration_tests/feature_tests/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..1b92138e0e --- /dev/null +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py @@ -0,0 +1,147 @@ +# 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 typing + +import numpy as np +import pytest + +import gt4py.next as gtx +from gt4py.next import Dims, Field, common, constructors, neighbor_sum + +from next_tests import definitions as test_defs +from next_tests.integration_tests import cases, cases_utils +from next_tests.integration_tests.cases_utils import ( # noqa: F401 [unused-import] # fixture + exec_alloc_descriptor, +) + + +class V(gtx.DimensionIndex): ... + + +class E(gtx.DimensionIndex): ... + + +class V2E(gtx.NeighborConnectivity[V, E], max_neighbors=4, min_neighbors=4): + class Local(gtx.LocalDimensionIndex): ... + + +#: A second connectivity over the same neighbor axis, bound to a different table. +class V2EShared(gtx.NeighborConnectivity[V, E]): + Local: typing.TypeAlias = V2E.Local + + +@pytest.fixture +def case(exec_alloc_descriptor): + mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) + v2e_arr = mesh.offset_provider[cases_utils.V2EDim.tag].asnumpy() + table = constructors.as_connectivity( + domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, + codomain=E, + data=v2e_arr, + skip_value=None, + allocator=exec_alloc_descriptor.allocator, + ) + common.check_neighbor_table(V2E, table) + shared_table = constructors.as_connectivity( + domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, + codomain=E, + data=np.ascontiguousarray(v2e_arr[:, ::-1]), + skip_value=None, + allocator=exec_alloc_descriptor.allocator, + ) + common.check_neighbor_table(V2EShared, shared_table) + return cases.Case( + ( + None + if isinstance(exec_alloc_descriptor, test_defs.EmbeddedDummyBackend) + else exec_alloc_descriptor + ), + # NOTE: still keyed on the local dimension's tag; class keys come with the removal of + # `FieldOffset`. + offset_provider={V2E.offset_tag: table, V2EShared.offset_tag: shared_table}, + default_sizes={V: mesh.num_vertices, E: mesh.num_edges, V2E.Local: v2e_arr.shape[1]}, + grid_type=common.GridType.UNSTRUCTURED, + allocator=exec_alloc_descriptor.allocator, + ) + + +def _table(case: cases.Case, connectivity=V2E) -> np.ndarray: + return case.offset_provider[connectivity.offset_tag].asnumpy() + + +@pytest.mark.uses_unstructured_shift +def test_shift(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2E[1]) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case)[:, 1]]) + + +@pytest.mark.uses_unstructured_shift +def test_reduction(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return neighbor_sum(a(V2E), axis=V2E.Local) + + cases.verify_with_default_data(case, testee, lambda a: np.sum(a[_table(case)], axis=1)) + + +@pytest.mark.uses_unstructured_shift +def test_sparse_argument(case): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + return neighbor_sum(s * a(V2E), axis=V2E.Local) + + cases.verify_with_default_data(case, testee, lambda s, a: np.sum(s * a[_table(case)], axis=1)) + + +@pytest.mark.uses_unstructured_shift +def test_program(case): + @gtx.field_operator + def shift_by_one(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2E[0]) + + @gtx.program + def testee(a: Field[Dims[E], float], out: Field[Dims[V], float]): + shift_by_one(a, out=out) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case)[:, 0]]) + + +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_offset_tag_differing_from_local_dim +def test_shift_through_a_shared_local_dimension(case): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2EShared[1]) + + cases.verify_with_default_data(case, testee, lambda a: a[_table(case, V2EShared)[:, 1]]) + + +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_offset_tag_differing_from_local_dim +@pytest.mark.uses_offset_tag_differing_from_local_dim_in_reduction +def test_reduction_through_a_shared_local_dimension(case): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + # combines a sparse field on the shared axis with each connectivity's neighbors + return neighbor_sum(s * a(V2EShared) - a(V2E), axis=V2E.Local) + + cases.verify_with_default_data( + case, + testee, + lambda s, a: np.sum(s * a[_table(case, V2EShared)] - a[_table(case)], axis=1), + ) diff --git a/tests/next_tests/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/iterator_tests/test_builtins.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py index b16461447a..094686be80 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_builtins.py @@ -59,7 +59,7 @@ class Node(gtx.DimensionIndex): ... -class NeighDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class NeighDim(gtx.LocalDimensionIndex): ... def array_maker(*lists): diff --git a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py index bd675b5f51..626e83d075 100644 --- a/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py +++ b/tests/next_tests/integration_tests/feature_tests/iterator_tests/test_strided_offset_provider.py @@ -17,7 +17,7 @@ from gt4py.next.iterator.embedded import StridedConnectivityField -class Dummy(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class Dummy(gtx.LocalDimensionIndex): ... class LocA(gtx.DimensionIndex): ... diff --git a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py index 4fbb01c72d..7051170dd7 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py @@ -39,11 +39,7 @@ # NOTE: imported, not redeclared. Under nominal identity (ADR 0028) a same-named declaration # here would be a different dimension from the one `toy_connectivity` declares, where the old # `Dimension("...")` values compared equal -- and tests mix objects from both modules. -from next_tests.toy_connectivity import E2VDim, Edge, V2EDim, Vertex - - -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) +from next_tests.toy_connectivity import E2V, E2VDim, Edge, V2E, V2EDim, Vertex def assert_close(expected, actual): diff --git a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py index 689d0d5f71..3cf6519d27 100644 --- a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py +++ b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py @@ -43,14 +43,14 @@ class E(gtx.DimensionIndex): ... #: N1 == N3 == N4, but N2 differs: the tag is `TaggedOffDim.tag`, the variable is `off_a`. -class TaggedOffDim(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class TaggedOffDim(gtx.LocalDimensionIndex): ... off_a = gtx.FieldOffset(TaggedOffDim.tag, source=E, target=(V, TaggedOffDim)) #: N1 == N2 == N4, but N3 differs: the local dimension is `Neigh`, the tag is `OffB`. -class Neigh(gtx.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Neigh(gtx.LocalDimensionIndex): ... OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) diff --git a/tests/next_tests/toy_connectivity.py b/tests/next_tests/toy_connectivity.py index c368b04245..28fffdc732 100644 --- a/tests/next_tests/toy_connectivity.py +++ b/tests/next_tests/toy_connectivity.py @@ -21,22 +21,26 @@ class Edge(gtx.DimensionIndex): ... class Cell(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2V(gtx.NeighborConnectivity[Edge, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class V2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2V(gtx.NeighborConnectivity[Vertex, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) -E2V = gtx.FieldOffset(E2VDim.tag, source=Vertex, target=(Edge, E2VDim)) -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) -V2V = gtx.FieldOffset(V2VDim.tag, source=Vertex, target=(Vertex, V2VDim)) +V2EDim = V2E.Local +E2VDim = E2V.Local +C2EDim = C2E.Local +V2VDim = V2V.Local # 3x3 periodic edges cells # 0 - 1 - 2 - 0 1 2 diff --git a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py index bef6a29c8a..3acc165129 100644 --- a/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py +++ b/tests/next_tests/unit_tests/embedded_tests/test_nd_array_field.py @@ -20,6 +20,7 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Field, DimensionIndex, @@ -49,10 +50,10 @@ class V(DimensionIndex): ... class E(DimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class V2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2V(LocalDimensionIndex): ... class C(DimensionIndex): ... @@ -61,7 +62,7 @@ class C(DimensionIndex): ... class K(DimensionIndex): ... -class C2E2CO(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E2CO(LocalDimensionIndex): ... class A(DimensionIndex): ... @@ -76,7 +77,7 @@ class X(DimensionIndex): ... class Y(DimensionIndex): ... -class L(DimensionIndex, kind=DimensionKind.LOCAL): ... +class L(LocalDimensionIndex): ... class S(DimensionIndex): ... @@ -97,7 +98,7 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2EDim(LocalDimensionIndex): ... class C2V(DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py index 25284281ef..9cdb145fca 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py @@ -6,6 +6,8 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import typing + import pytest import gt4py.next as gtx @@ -21,11 +23,14 @@ class VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... class Dim(gtx.DimensionIndex): ... -class LocalDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class LocalDim(gtx.LocalDimensionIndex): ... CartesianOffset = gtx.FieldOffset("CartesianOffset", source=Dim, target=(Dim,)) -UnstructuredOffset = gtx.FieldOffset(LocalDim.tag, source=Dim, target=(Dim, LocalDim)) + + +class UnstructuredOffset(gtx.NeighborConnectivity[Dim, Dim]): + Local: typing.TypeAlias = LocalDim def test_domain_deduction_cartesian(): diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index f696daf1b4..65a421c171 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py @@ -49,10 +49,11 @@ class Edge(gtx.DimensionIndex): ... class Vertex(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... -V2E = gtx.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) +V2EDim = V2E.Local class TDim(gtx.DimensionIndex): ... @@ -63,7 +64,7 @@ class TDim(gtx.DimensionIndex): ... #: An offset whose tag differs from the name of the Python variable it is bound to, and #: from the name of its local dimension. Lowering must emit the *tag*. -class RenamedV2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class RenamedV2EDim(gtx.LocalDimensionIndex): ... renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py index 290d2914fc..7051400ed8 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_source_utils.py @@ -19,7 +19,7 @@ """ import gt4py.next as gtx -from gt4py.next import Dims, Dimension, DimensionIndex, float64, neighbor_sum +from gt4py.next import Dims, Dimension, DimensionIndex, LocalDimensionIndex, float64, neighbor_sum from gt4py.next.ffront import source_utils from gt4py.next.ffront.source_utils import get_closure_vars_from_function @@ -30,10 +30,11 @@ class Cell(DimensionIndex): ... class Edge(DimensionIndex): ... -class C2EDim(DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2EDim = C2E.Local CField = gtx.Field[Dims[Cell], float64] EField = gtx.Field[Dims[Edge], float64] diff --git a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py index 88a8c640d8..c21b88be36 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_type_deduction.py @@ -18,6 +18,8 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, Field, FieldOffset, astype, @@ -50,7 +52,11 @@ class X(DimensionIndex): ... class Y(DimensionIndex): ... -class Y2XDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class Y2X(NeighborConnectivity[Y, X]): + class Local(LocalDimensionIndex): ... + + +Y2XDim = Y2X.Local class K(DimensionIndex, kind=DimensionKind.VERTICAL): ... @@ -71,7 +77,11 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + + +V2EDim = V2E.Local class IDim(DimensionIndex): ... @@ -286,7 +296,6 @@ def domain_comparison(a: Field[[TDim], float], b: Field[[TDim], float]): @pytest.fixture def premap_setup(): - Y2X = FieldOffset(Y2XDim.tag, source=X, target=(Y, Y2XDim)) return X, Y, Y2XDim, Y2X @@ -567,8 +576,6 @@ def as_offset_dtype(a: Field[[ADim, BDim], float], b: Field[[BDim], float]): def test_as_offset_non_cartesian(): - V2E = FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) - def as_offset_neighbor(a: Field[[Edge], float], b: Field[[Edge], int]): return a(as_offset(V2E, b)) diff --git a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py index 7120882c75..793f802c9f 100644 --- a/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py +++ b/tests/next_tests/unit_tests/iterator_tests/ir_utils_tests/test_domain_utils.py @@ -33,13 +33,13 @@ class Vertex(common.DimensionIndex): ... class Edge(common.DimensionIndex): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... -class V2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2VDim(common.LocalDimensionIndex): ... a_range = domain_utils.SymbolicRange(0, 10) diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py index 12c0b75649..be0048c285 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py @@ -29,10 +29,11 @@ class E(gtx.DimensionIndex): ... class V(gtx.DimensionIndex): ... -class E2VDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class E2V(gtx.NeighborConnectivity[E, V]): + class Local(gtx.LocalDimensionIndex): ... -E2V = gtx.FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) +E2VDim = E2V.Local # 0 --0-- 1 --1-- 2 diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py index eafe45b9d3..23c4dbaf07 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_domain_inference.py @@ -45,7 +45,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class E2VDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class E2VDim(common.LocalDimensionIndex): ... float_i_field = ts.FieldType(dims=[IDim], dtype=float_type) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py index 5bceba0e53..7dfbf60934 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_fuse_as_fieldop.py @@ -17,7 +17,7 @@ from gt4py.next.type_system import type_specifications as ts -class Neighbor(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Neighbor(common.LocalDimensionIndex): ... class IDim(common.DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py index 3f6cd212ae..d58b728e10 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_empty_concat_where.py @@ -25,7 +25,7 @@ class Vertex(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class Edge(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... -class V2EDim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class V2EDim(common.LocalDimensionIndex): ... class K(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py index 5451a2dc44..3f93017c67 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py @@ -27,10 +27,10 @@ class dummy_neighbor(common.DimensionIndex): ... #: The local dimensions of the neighbor lists under test. Each one's `tag` is also its IR offset #: string and its offset-provider key: `UnrollReduce` looks a connectivity up by the local #: dimension of the list it reduces, so those three names must be a single string (ADR 0028). -class Dim(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim(common.LocalDimensionIndex): ... -class Dim2(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class Dim2(common.LocalDimensionIndex): ... def dummy_connectivity_type(max_neighbors: int, has_skip_values: bool): diff --git a/tests/next_tests/unit_tests/otf_tests/test_runners.py b/tests/next_tests/unit_tests/otf_tests/test_runners.py index 2ea966ff7e..f75471fcda 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_runners.py +++ b/tests/next_tests/unit_tests/otf_tests/test_runners.py @@ -31,7 +31,7 @@ class Vertex(gtx.DimensionIndex): ... class Edge(gtx.DimensionIndex): ... -class V2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ... +class V2EDim(gtx.LocalDimensionIndex): ... @pytest.fixture diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index 1bec200fad..873a0f5bb4 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -94,7 +94,7 @@ def test_find_constant_symbols(has_unit_stride, disable_field_origin): itir.SetAt( expr=im.as_fieldop( im.lambda_("it")(im.reduce("plus", im.literal_from_value(1.0))(im.deref("it"))) - )(im.as_fieldop_neighbors(V2E.value, "x")), + )(im.as_fieldop_neighbors(V2E.Local.tag, "x")), domain=im.get_field_domain(gtx_common.GridType.UNSTRUCTURED, "y", VFTYPE.dims), target=itir.SymRef(id="y"), ) diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index 500a355dbb..d1b9f97560 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -22,6 +22,7 @@ Dimension, DimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Infinity, UnitRange, @@ -57,19 +58,19 @@ class I(common.DimensionIndex): ... class I_half(common.DimensionIndex): ... -class C2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class C2E(LocalDimensionIndex): ... -class V2E(DimensionIndex, kind=DimensionKind.LOCAL): ... +class V2E(LocalDimensionIndex): ... -class E2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2V(LocalDimensionIndex): ... -class E2C(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C(LocalDimensionIndex): ... -class E2C2V(DimensionIndex, kind=DimensionKind.LOCAL): ... +class E2C2V(LocalDimensionIndex): ... class ECDim(DimensionIndex): ... diff --git a/tests/next_tests/unit_tests/test_custom_layout_allocators.py b/tests/next_tests/unit_tests/test_custom_layout_allocators.py index f5257506fb..53312af3ab 100644 --- a/tests/next_tests/unit_tests/test_custom_layout_allocators.py +++ b/tests/next_tests/unit_tests/test_custom_layout_allocators.py @@ -29,13 +29,13 @@ class D2(common.DimensionIndex, kind=common.DimensionKind.HORIZONTAL): ... class D0_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... -class D1_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D1_local(common.LocalDimensionIndex): ... class D2_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... -class D2_local(common.DimensionIndex, kind=common.DimensionKind.LOCAL): ... +class D2_local(common.LocalDimensionIndex): ... class D1_vertical(common.DimensionIndex, kind=common.DimensionKind.VERTICAL): ... diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py new file mode 100644 index 0000000000..86bb41f1a3 --- /dev/null +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -0,0 +1,512 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import pickle +import textwrap +import typing + +import numpy as np +import pytest + +from gt4py._core import definitions as core_defs +from gt4py.next import common +from gt4py.next.common import ( + DimensionIndex, + DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, + NeighborConnectivityType, +) +from gt4py.next.ffront import transform_utils +from gt4py.next.type_system import type_specifications as ts, type_translation + + +class Vertex(DimensionIndex): ... + + +class Edge(DimensionIndex): ... + + +class KDim(DimensionIndex, kind=DimensionKind.VERTICAL): ... + + +class V2E(NeighborConnectivity[Vertex, Edge], max_neighbors=4, min_neighbors=3): + class Local(LocalDimensionIndex): ... + + +class E2V(NeighborConnectivity[Edge, Vertex]): + class Local(LocalDimensionIndex): ... + + +class LsqCoeff(LocalDimensionIndex, size=3): ... + + +def _declare(source: str) -> dict: + """ + Run `source` as the body of a throwaway module. + + Declarations have to be at module level, so error cases cannot simply be written inside + the test function: the `` check would fire before the one under test. + """ + namespace = { + "__name__": __name__, + "typing": typing, + "DimensionIndex": DimensionIndex, + "DimensionKind": DimensionKind, + "LocalDimensionIndex": LocalDimensionIndex, + "NeighborConnectivity": NeighborConnectivity, + "Vertex": Vertex, + "Edge": Edge, + "KDim": KDim, + "V2E": V2E, + "ConstListDim": common.ConstListDim, + } + 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 + """, + "counts are declared by its owner", + ), + ( + """ + 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_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_field_offset_is_derived_once(self): + assert V2E.__gt_field_offset__() is V2E.__gt_field_offset__() + assert V2E.__gt_field_offset__().value == V2E.Local.tag + + def test_neighbor_index_accepts_numpy_integers(self): + from gt4py.next import constructors, embedded + + table = constructors.as_connectivity( + domain={Vertex: 2, V2E.Local: 4}, + codomain=Edge, + data=np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), + ) + with embedded.context.update(offset_provider={V2E.Local.tag: table}): + assert np.array_equal(V2E[np.int32(1)].asnumpy(), V2E[1].asnumpy()) + + def test_legacy_field_offset_has_local(self): + from gt4py.next import FieldOffset + + assert FieldOffset("V2E", source=Edge, target=(Vertex, V2E.Local)).Local is V2E.Local + + def test_attribute_errors_are_dsl_errors(self): + from gt4py.next import errors, field_operator + from gt4py.next.ffront.func_to_foast import FieldOperatorParser + from gt4py.next import Dims, Field + + def origin_of(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: + return a(V2E.origin) + + with pytest.raises(errors.DSLError, match="has no attribute 'origin'"): + FieldOperatorParser.apply_to_function(origin_of) + + def test_fingerprint_covers_the_declaration(self): + from gt4py.next import fingerprinting + + def fingerprint_of(source: str) -> str: + # lenient: `_declare` classes are not importable, as in a re-run notebook cell + return fingerprinting.lenient_fingerprinter(_declare(source)["C"]) + + base = fingerprint_of( + """ + class C(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + """ + ) + swapped = fingerprint_of( + """ + class C(NeighborConnectivity[Edge, Vertex]): + class Local(LocalDimensionIndex): ... + """ + ) + counted = fingerprint_of( + """ + class C(NeighborConnectivity[Vertex, Edge], max_neighbors=3): + class Local(LocalDimensionIndex): ... + """ + ) + assert len({base, swapped, counted}) == 3 + assert fingerprinting.strict_fingerprinter(V2E) != fingerprinting.strict_fingerprinter(E2V) + + def test_grid_type_deduction(self): + assert ( + transform_utils._deduce_grid_type(None, [Vertex, V2E]) is common.GridType.UNSTRUCTURED + ) + with pytest.raises(ValueError, match="CARTESIAN"): + transform_utils._deduce_grid_type(common.GridType.CARTESIAN, [V2E]) + + +def test_redefined_declaration_with_an_adopted_local(monkeypatch): + """Re-running a cell must re-own the adopted local dimension, not become a sharer.""" + import sys + import types as pytypes + + module = pytypes.ModuleType("_readopted_connectivity_module") + monkeypatch.setitem(sys.modules, module.__name__, module) + source = textwrap.dedent( + """ + import typing + + from gt4py.next.common import DimensionIndex, LocalDimensionIndex, NeighborConnectivity + + class V(DimensionIndex): ... + class E(DimensionIndex): ... + class V2EDim(LocalDimensionIndex, size={n}): ... + class V2E(NeighborConnectivity[V, E], max_neighbors={n}): + Local: typing.TypeAlias = V2EDim + """ + ) + exec(source.format(n=4), module.__dict__) + assert module.V2E.offset_tag == module.V2EDim.tag + + # the redefinition takes ownership over again, and its counts are checked against `size=` + exec(source.format(n=2), module.__dict__) + assert module.V2EDim.owner is module.V2E + assert module.V2E.offset_tag == module.V2EDim.tag + assert module.V2EDim.max_neighbors == 2 + + +def test_local_dimension_of(): + shared = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + )["V2EShared"] + assert common.local_dimension_of(V2E) is V2E.Local + assert common.local_dimension_of(shared) is V2E.Local + with pytest.raises(TypeError, match="not a connectivity declaration"): + common.local_dimension_of(NeighborConnectivity) + + +class TestFieldOffsetDeprecation: + def test_unstructured_field_offset_warns(self): + from gt4py.next import FieldOffset + + with pytest.warns(DeprecationWarning, match="NeighborConnectivity"): + FieldOffset(V2E.Local.tag, source=Edge, target=(Vertex, V2E.Local)) + + def test_derived_and_cartesian_field_offsets_do_not_warn(self, recwarn): + from gt4py.next import FieldOffset + + class_ns = _declare( + """ + class C2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + """ + ) + class_ns["C2E"].__gt_field_offset__() + FieldOffset("Koff", source=KDim, target=(KDim,)) + assert not [w for w in recwarn if issubclass(w.category, DeprecationWarning)] + + +def test_the_const_list_dimension_cannot_be_adopted(): + with pytest.raises(TypeError, match="cannot adopt"): + _declare( + """ + class C(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = ConstListDim + """ + ) 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..45ddb1037f 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] | FieldOffset | type[NeighborConnectivity[Any, Any]], *args: Connectivity[Any, Any] | FieldOffset | type[NeighborConnectivity[Any, Any]]) -> Field[Any, Any]" + + - case: neighbor_connectivity_generic_local + main: | + from __future__ import annotations + import typing + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + L = typing.TypeVar("L", bound=gtx.LocalDimensionIndex) + + def local_of(conn: type[gtx.NeighborConnectivity]) -> type[gtx.LocalDimensionIndex]: + # `conn.Local` is not annotated, so that a declaration's `Local` stays a type; generic + # code reads it through the accessor + return gtx.local_dimension_of(conn) + + def first(a: gtx.Field[gtx.Dims[Vertex, L], gtx.float64]) -> type[L]: + raise NotImplementedError + + def caller(a: gtx.Field[gtx.Dims[Vertex, V2E.Local], gtx.float64]) -> None: + reveal_type(first(a)) + out: | + main:22:17: note: Revealed type is "type[main.V2E.Local]" diff --git a/uv.lock b/uv.lock index 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"