diff --git a/docs/development/ADRs/next/0019-Connectivities.md b/docs/development/ADRs/next/0019-Connectivities.md index 0c80a6cf55..5357c9f078 100644 --- a/docs/development/ADRs/next/0019-Connectivities.md +++ b/docs/development/ADRs/next/0019-Connectivities.md @@ -9,6 +9,10 @@ tags: [] - **Created**: 2024-11-08 - **Updated**: 2026-05-27 +> The `FieldOffset` part of this record is superseded by +> [ADR 0030](0030-Connectivities_As_Types.md): connectivities are declared as +> `NeighborConnectivity` classes, and offset providers are keyed by them. + The representation of Connectivities (neighbor tables, `NeighborTableOffsetProvider`) and their identifier (offset tag, `FieldOffset`, etc.) was extended and modified based on the needs of different parts of the toolchain. Here we outline the ideas for consolidating the different closely-related concepts. ## History diff --git a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md index fcb29cbd04..03458b9b7f 100644 --- a/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md +++ b/docs/development/ADRs/next/0029-Dimensions_As_Nominal_Types.md @@ -53,7 +53,8 @@ string equality and are never checked against each other at declaration time: th name, and the `offset_provider` key. Whichever one reaches `common.get_offset` depends on the execution path and the operation. Making a dimension's identity its Python type is the prerequisite for collapsing those -names into one declaration (a follow-up ADR covers the connectivity half). +names into one declaration ([ADR 0030](0030-Connectivities_As_Types.md) covers the +connectivity half). ## Decision @@ -82,10 +83,11 @@ disappears. 1. **Reconstruction from the IR is an import.** `common.resolve(tag)` imports the module and walks the qualname; nested declarations resolve naturally. The IR - references a Python type exactly the way `pickle` references a class. It is - memoized, because type inference calls it once per `AxisLiteral`. An - `AxisLiteral` stores only the tag: its `kind` is the resolved dimension's, so the - two cannot disagree. + references a Python type exactly the way `pickle` references a class. Where the + module path ends is memoized, because type inference calls it once per + `AxisLiteral`; the attribute walk is repeated, so a redefined declaration is + found. An `AxisLiteral` stores only the tag: its `kind` is the resolved + dimension's, so the two cannot disagree. A purely dotted tag does not record where the module path ends and the qualname begins, so `resolve` tries the *longest importable prefix* and walks @@ -216,7 +218,8 @@ The last row needs `DimensionMeta.__add__` / `__sub__` declared with the self-ty site and both reject it at the definition site, with different diagnostics (mypy `[misc]`, pyright `reportGeneralTypeIssues`), so it costs two separately spelled suppressions. The runtime check covers unannotated code; hand-written iterator IR, -which names dimensions by tag, is not checked. Comparisons are deliberately *not* restricted: `D == n` +which names dimensions by tag, is not checked. `as_offset(dim, field)` needs index +arithmetic too and takes an `AnyCartesianAxisIndex`. Comparisons are deliberately *not* restricted: `D == n` and `D < n` build a `Domain` on every dimension, as `concat_where` over a mesh location requires. diff --git a/docs/development/ADRs/next/0030-Connectivities_As_Types.md b/docs/development/ADRs/next/0030-Connectivities_As_Types.md index 07ea6992b4..a5c88f32b5 100644 --- a/docs/development/ADRs/next/0030-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0030-Connectivities_As_Types.md @@ -75,7 +75,8 @@ and skip values were never checked against the `FieldOffset` declaration. `(Domain, V2E.Local)`, the codomain is `Codomain`, the dtype is integral, and the neighbor counts and skip values agree. Skip values are checked on the table's type: a table with a `skip_value` counts as having skip values whether - or not an entry uses it. + or not an entry uses it. Programs run the check on the tables they are given, + see below. `Domain` and `Codomain` name the two index spaces the declaration maps between. A bound table is a field over `(Domain, Local)` with values in `Codomain`: the @@ -93,8 +94,7 @@ of a table bound to a declaration is a `common.NeighborTableType`: `domain` and `codomain` are derived from the declaration, `(connectivity.domain, local_dimension_of(connectivity))` and `connectivity.codomain`, so they cannot disagree with it. The mapping from -offset-provider keys to these records is `common.TableTypes`, and it can be -given instead of the tables for ahead-of-time compilation. +offset-provider keys to these records is `common.TableTypes`. A table cannot tell which declaration it is bound to: the table of a sharer (`C2CE`) has the same domain as its owner's (`C2E`), with another codomain. So a @@ -173,11 +173,11 @@ treating it as a type. ### Frontend integration -A declaration is typed like the `FieldOffset` it replaces: `V2E.__gt_type__()` -is a `ts.ShiftType`, which takes a field over the codomain to one over the -domain, `Shift[: Edge -> (Vertex, V2E.Local)]`. `V2E[i]` has the domain -`(Vertex,)`, and so does a Cartesian shift `KDim + 1`, over `KDim` and without a -tag. The tag is the connectivity's `offset_tag`: +A declaration is typed as a shift: `V2E.__gt_type__()` is a `ts.ShiftType`, +which takes a field over the codomain to one over the domain, +`Shift[: Edge -> (Vertex, V2E.Local)]`. `V2E[i]` has the domain `(Vertex,)`, +and so do the Cartesian shifts `KDim + 1` and `as_offset(KDim, offsets)`, over +`KDim` and without a tag. The tag is the connectivity's `offset_tag`: - **the local dimension's tag**, `V2E.Local.tag`, for the connectivity that declares it. This is the single string that shifts, neighbor reductions and @@ -190,15 +190,64 @@ tag. The tag is the connectivity's `offset_tag`: over it (`common.connectivity_key_over`): the owner's if bound, else the sharer with the smallest tag. Connectivities sharing a local dimension must therefore have the same neighbor *structure* — the same count, and a skip value - at the same positions — which is what sharing a neighbor axis means. + at the same positions — which is what sharing a neighbor axis means; + `check_offset_provider` enforces it for the tables it is given. -`V2E.Local` inside DSL code types as that local dimension, 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 +`V2E.Local` inside DSL code types as that local dimension. The other frontend +touch points treat a declaration as an unstructured shift: grid-type deduction +(`transform_utils`, `past_to_itir`) counts it as unstructured, and embedded +`premap` accepts it. `V2E[i]` subscripts the metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it. +### Offset providers are keyed by the declaration + +Users bind tables to declarations: + +```python +program(..., offset_provider={V2E: v2e_table, C2E: c2e_table}) +``` + +Every entry point of a program normalizes such a provider to the form the IR +uses: each declaration is replaced by its `offset_tag`. Everything below the +entry points — lowering, the backends, compiled-program caching — therefore keeps +seeing a provider keyed by strings, which is also what hand-written IR uses. + +The frontend entry points (`Program.__call__`, `FieldOperator.__call__`, +`compile`, `CompilationOptions.connectivities`) are *strict*: a string key must +be a tag, i.e. a qualified name, and a bare name such as `"V2E"` is the removed +`FieldOffset` spelling, rejected with a message pointing here. The IR-level hooks +(`embedded.context.update`, the iterator `fendef`, DaCe's `get_sdfg_conn_args`) +accept any string, because a hand-written program names its offsets itself. + +The types of the tables, a `common.TableTypes` under the same keys, are called +`table_types` throughout: `CompileTimeArgs.table_types`, the IR passes and the code +generators. Ahead-of-time compilation can take them in place of the tables, +e.g. `{V2E: NeighborTableType(connectivity=V2E, dtype=int32, skip_value=None, max_neighbors=6)}` (`compile(offset_provider=...)` accepts either), and they are +keyed and normalized exactly like an offset provider: by declarations, strictly, +at the frontend; by tags at the IR level. + +Tables are checked against their declarations (`check_offset_provider`) at every +entry point, but the result is remembered per set of bound tables, so repeated +calls cost one hash. Reading the tables — comparing the skip-value positions of +two connectivities that share a local dimension — is done only where a program is +compiled, not on the call path. A tag that +names no declared connectivity, as in hand-written IR, is not checked. + +### `FieldOffset` is removed + +`FieldOffset` and its export are gone. An unstructured connectivity is a +`NeighborConnectivity`; a Cartesian shift is `Dim + i`, which the DSL already +had; and `as_offset` takes the Cartesian axis to shift along, +`as_offset(KDim, k_offsets)`, instead of a Cartesian `FieldOffset`. +`scripts/python/migrate_connectivities.py` rewrites declarations and Cartesian +offset uses, and reports the provider keys and other sites it cannot rewrite from +the source alone. It decides per declaration whether a dimension is a Cartesian +axis (`kind=VERTICAL`, a Cartesian offset, `as_offset`, index arithmetic or a +staggered counterpart) or a mesh location (a domain or codomain of a neighbor +offset), and reports a dimension with no evidence, or with evidence for both, +instead of guessing. It drops aliases of the removed `DimensionKind.LOCAL` and +reports its other uses. + ## Consequences - An unstructured connectivity is spelled once. The provider key, the offset tag @@ -211,8 +260,8 @@ metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, - 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. +- `FieldOffset` is removed, and offset providers are keyed by declarations: a + breaking change for every unstructured program, eased by the migration script. ## Alternatives considered diff --git a/docs/user/next/QuickstartGuide.md b/docs/user/next/QuickstartGuide.md index 1695c0a63f..9357effbeb 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.LocalDimensionIndex): ... -E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim)) +class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + class Local(gtx.LocalDimensionIndex): ... ``` -The field offset is named by its local dimension's `tag`, and the offset provider below is keyed by the same `tag`, so all three refer to one connectivity. Note that the field offset does not contain the actual connectivity table, that's provided through an _offset provider_: +Note that the declaration does not contain the actual connectivity table, that's provided through an _offset provider_, a dictionary from connectivity declarations to tables: ```{code-cell} ipython3 -E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2CDim], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) +E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2C.Local], codomain=CellDim, data=edge_to_cell_table, skip_value=-1) ``` The field operator `nearest_cell_to_edge` below shows an example of applying this transform. There is a little twist though: the subscript in `E2C[0]` means that only the value of the first connected cell is taken, the second (if exists) is ignored. -Pay attention to the syntax where the field offset `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. +Pay attention to the syntax where the connectivity `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program. ```{code-cell} ipython3 @gtx.field_operator @@ -250,7 +250,7 @@ def nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64]) -> gtx. def run_nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): nearest_cell_to_edge(cell_values, out=out) -run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2C: E2C_offset_provider}) print("0th adjacent cell's value: {}".format(edge_values.asnumpy())) ``` @@ -265,19 +265,19 @@ Running the above snippet results in the following edge field: #### Using reductions on connected mesh elements -Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2CDim) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2CDim` dimension. +Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2C.Local) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2C.Local` dimension. ```{code-cell} ipython3 @gtx.field_operator def sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64]) -> gtx.Field[Dims[EdgeDim], float64]: - # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2CDim], float64] - return neighbor_sum(cells(E2C), axis=E2CDim) + # type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2C.Local], float64] + return neighbor_sum(cells(E2C), axis=E2C.Local) @gtx.program def run_sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]): sum_adjacent_cells(cells, out=out) -run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider}) +run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2C: E2C_offset_provider}) print("sum of adjacent cells: {}".format(edge_values.asnumpy())) ``` @@ -376,13 +376,13 @@ print("where nested tuple return: {}".format(((result_1.asnumpy(), result_2.asnu #### Implementing the pseudo-laplacian -As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to define the local dimension, the field offset, and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: +As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to declare the connectivity and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before: ```{code-cell} ipython3 -class C2EDim(gtx.LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=EdgeDim, target=(CellDim, C2EDim)) +class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + class Local(gtx.LocalDimensionIndex): ... -C2E_offset_provider = gtx.as_connectivity([CellDim, C2EDim], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) +C2E_offset_provider = gtx.as_connectivity([CellDim, C2E.Local], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1) ``` **Weights of edge differences:** @@ -410,7 +410,7 @@ edge_weights = np.array([ [0, -1, -1], # cell 5 ], dtype=np.float64) -edge_weight_field = gtx.as_field([CellDim, C2EDim], edge_weights) +edge_weight_field = gtx.as_field([CellDim, C2E.Local], edge_weights) ``` Now you have everything to implement the pseudo-laplacian. Its field operator requires the cell field and the edge weights as inputs, and outputs a cell field of the same shape as the input. @@ -422,9 +422,9 @@ The second lines first creates a temporary field using `edge_differences(C2E)`, ```{code-cell} ipython3 @gtx.field_operator def pseudo_lap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: edges = cells(E2C[0]) # type: gtx.Field[Dims[EdgeDim], float64] - return neighbor_sum(edges(C2E) * edge_weights, axis=C2EDim) + return neighbor_sum(edges(C2E) * edge_weights, axis=C2E.Local) ``` The program itself is just a shallow wrapper over the `pseudo_lap` field operator. The significant part is how offset providers for both the edge-to-cell and cell-to-edge connectivities are supplied when the program is called: @@ -432,7 +432,7 @@ The program itself is just a shallow wrapper over the `pseudo_lap` field operato ```{code-cell} ipython3 @gtx.program def run_pseudo_laplacian(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64], + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64], out : gtx.Field[Dims[CellDim], float64]): pseudo_lap(cells, edge_weights, out=out) @@ -441,7 +441,7 @@ result_pseudo_lap = gtx.as_field([CellDim], np.zeros(shape=(6,))) run_pseudo_laplacian(cell_values, edge_weight_field, result_pseudo_lap, - offset_provider={E2CDim.tag: E2C_offset_provider, C2EDim.tag: C2E_offset_provider}) + offset_provider={E2C: E2C_offset_provider, C2E: C2E_offset_provider}) print("pseudo-laplacian: {}".format(result_pseudo_lap.asnumpy())) ``` @@ -451,7 +451,7 @@ As a closure, here is an example of chaining field operators, which is very simp ```{code-cell} ipython3 @gtx.field_operator def pseudo_laplap(cells : gtx.Field[Dims[CellDim], float64], - edge_weights : gtx.Field[Dims[CellDim, C2EDim], float64]) -> gtx.Field[Dims[CellDim], float64]: + edge_weights : gtx.Field[Dims[CellDim, C2E.Local], float64]) -> gtx.Field[Dims[CellDim], float64]: return pseudo_lap(pseudo_lap(cells, edge_weights), edge_weights) ``` diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb index b0a1980d0f..0d98021f37 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise.ipynb @@ -126,7 +126,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb index 573ee6a44e..5176f48285 100644 --- a/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/2_divergence_exercise_solution.ipynb @@ -131,7 +131,7 @@ " A,\n", " edge_orientation,\n", " out=divergence_gt4py,\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb index 2b422b1823..77d1ecf296 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise.ipynb @@ -123,7 +123,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb index 85044b989f..b90958ea63 100644 --- a/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/3_gradient_exercise_solution.ipynb @@ -136,7 +136,7 @@ " A,\n", " edge_orientation,\n", " out=(gradient_gt4py_x, gradient_gt4py_y),\n", - " offset_provider={C2E.value: c2e_connectivity},\n", + " offset_provider={C2E: c2e_connectivity},\n", " )\n", "\n", " assert np.allclose(gradient_gt4py_x.asnumpy(), gradient_numpy_x)\n", diff --git a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb index dc321f1bdd..ed332054f8 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise.ipynb @@ -147,7 +147,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb index 251fe8239a..0f76b5a6b1 100644 --- a/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/4_curl_exercise_solution.ipynb @@ -152,7 +152,7 @@ " dualA,\n", " edge_orientation,\n", " out=curl_gt4py,\n", - " offset_provider={V2E.value: v2e_connectivity},\n", + " offset_provider={V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(curl_gt4py.asnumpy(), divergence_ref)" diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb index 30f568de6f..b557d404ea 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise.ipynb @@ -293,10 +293,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E: c2e_connectivity,\n", + " V2E: v2e_connectivity,\n", + " E2V: e2v_connectivity,\n", + " E2C: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb index eaeb8c7b02..7e1a7ecc4c 100644 --- a/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/5_vector_laplace_exercise_solution.ipynb @@ -314,10 +314,10 @@ " edge_orientation_cell,\n", " out=laplacian_gt4py,\n", " offset_provider={\n", - " C2E.value: c2e_connectivity,\n", - " V2E.value: v2e_connectivity,\n", - " E2V.value: e2v_connectivity,\n", - " E2C.value: e2c_connectivity,\n", + " C2E: c2e_connectivity,\n", + " V2E: v2e_connectivity,\n", + " E2V: e2v_connectivity,\n", + " E2C: e2c_connectivity,\n", " },\n", " )\n", "\n", diff --git a/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb b/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb index 5514c1b4f7..194d72bcc1 100644 --- a/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb +++ b/docs/user/next/workshop/exercises/6_where_domain_solution.ipynb @@ -326,23 +326,7 @@ "execution_count": 12, "id": "ed393959", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Test successful\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/var/folders/2b/_2y31vzs4sl_7rngh2yghbpw0000gn/T/ipykernel_56692/57164705.py:4: UserWarning: Field View Program 'program_domain_where': Using Python execution, consider selecting a perfomance backend.\n", - " program_domain_where(a, b, offset_provider={\"Koff\": K})\n" - ] - } - ], + "outputs": [], "source": [ "test_domain_where()\n", "print(\"Test successful\")" diff --git a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb index b278cee26d..26973ad5b4 100644 --- a/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb +++ b/docs/user/next/workshop/exercises/8_diffusion_exercise_solution.ipynb @@ -169,7 +169,7 @@ " kappa,\n", " dt,\n", " out=(divergence_gt4py_1, divergence_gt4py_2),\n", - " offset_provider={E2C2V.value: e2c2v_connectivity, V2E.value: v2e_connectivity},\n", + " offset_provider={E2C2V: e2c2v_connectivity, V2E: v2e_connectivity},\n", " )\n", "\n", " assert np.allclose(divergence_gt4py_1.asnumpy(), divergence_ref_1)\n", diff --git a/docs/user/next/workshop/exercises/helpers.py b/docs/user/next/workshop/exercises/helpers.py index 272fbdcc57..176b6a6ed5 100644 --- a/docs/user/next/workshop/exercises/helpers.py +++ b/docs/user/next/workshop/exercises/helpers.py @@ -15,9 +15,9 @@ CartesianAxisIndex, Dimension, DimensionIndex, - LocalDimensionIndex, DimensionKind, - FieldOffset, + LocalDimensionIndex, + NeighborConnectivity, ) from gt4py.next.program_processors.runners import roundtrip from gt4py.next.program_processors.runners.gtfn import ( @@ -396,31 +396,36 @@ class E(DimensionIndex): ... class K(CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... -class C2EDim(LocalDimensionIndex): ... +class C2E(NeighborConnectivity[C, E]): + class Local(LocalDimensionIndex): ... -C2E = FieldOffset(C2EDim.tag, source=E, target=(C, C2EDim)) +C2EDim = C2E.Local -class V2EDim(LocalDimensionIndex): ... +class V2E(NeighborConnectivity[V, E]): + class Local(LocalDimensionIndex): ... -V2E = FieldOffset(V2EDim.tag, source=E, target=(V, V2EDim)) +V2EDim = V2E.Local -class E2VDim(LocalDimensionIndex): ... +class E2V(NeighborConnectivity[E, V]): + class Local(LocalDimensionIndex): ... -E2V = FieldOffset(E2VDim.tag, source=V, target=(E, E2VDim)) +E2VDim = E2V.Local -class E2CDim(LocalDimensionIndex): ... +class E2C(NeighborConnectivity[E, C]): + class Local(LocalDimensionIndex): ... -E2C = FieldOffset(E2CDim.tag, source=C, target=(E, E2CDim)) +E2CDim = E2C.Local -class E2C2VDim(LocalDimensionIndex): ... +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 a336ed664a..aae419b52f 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.LocalDimensionIndex): ...\n", + "class E2C(gtx.NeighborConnectivity[Edge, Cell]):\n", + " class Local(gtx.LocalDimensionIndex): ...\n", "\n", "\n", - "E2C = gtx.FieldOffset(E2CDim.tag, source=Cell, target=(Edge, E2CDim))" + "E2CDim = E2C.Local" ] }, { @@ -320,7 +321,7 @@ " nearest_cell_to_edge(cell_field, out=edge_field)\n", "\n", "\n", - "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", + "run_nearest_cell_to_edge(cell_field, edge_field, offset_provider={E2C: E2C_offset_provider})\n", "\n", "print(\"0th adjacent cell's value: {}\".format(edge_field.asnumpy()))" ] @@ -395,7 +396,7 @@ " sum_adjacent_cells(cell_field, out=edge_field)\n", "\n", "\n", - "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={E2CDim.tag: E2C_offset_provider})\n", + "run_sum_adjacent_cells(cell_field, edge_field, offset_provider={E2C: E2C_offset_provider})\n", "\n", "print(\"sum of adjacent cells: {}\".format(edge_field.asnumpy()))" ] diff --git a/scripts/python/migrate_connectivities.py b/scripts/python/migrate_connectivities.py new file mode 100644 index 0000000000..1325a3b979 --- /dev/null +++ b/scripts/python/migrate_connectivities.py @@ -0,0 +1,682 @@ +#!/usr/bin/env -S uv run -q --frozen --isolated --python 3.12 --group scripts python3 +# +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +""" +Migrate gt4py.next user code to dimension and connectivity classes (ADRs 0029, 0030). + +Rewrites module-level declarations and the uses of Cartesian offsets: + + KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) + E2CDim = gtx.Dimension("E2C", gtx.DimensionKind.LOCAL) + E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim, E2CDim)) + Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) + ... a(Koff[1]) ... as_offset(Koff, k_field) ... + +becomes + + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... + class E2CDim(gtx.LocalDimensionIndex): ... + class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + Local: typing.TypeAlias = E2CDim + ... a(KDim + 1) ... as_offset(KDim, k_field) ... + +Whether a dimension is a Cartesian axis (`CartesianAxisIndex`: index arithmetic, a staggered +partner) or a mesh location (`DimensionIndex`) is decided per declaration, from the evidence in +all the given modules, in this order: + +1. a Cartesian axis if it is `kind=VERTICAL`, is the source of a Cartesian `FieldOffset`, is + used with `as_offset` or in index arithmetic (`K + 1`), or has a staggered counterpart + (`flip_staggered(K)`, or a `"_Staggered"` string under the old encoding); +2. a mesh location if it is the source or the target dimension of a neighbor `FieldOffset`; +3. otherwise a `DimensionIndex`, reported: nothing in the source tells a structured horizontal + axis from a mesh location. A dimension with evidence for both is reported as a conflict. + +A connectivity adopts its existing local dimension (`Local: TypeAlias = E2CDim`), so the names already used +for local dimensions keep working, and so do offsets that share a local dimension (`C2CE` +with `C2EDim`). The types of an offset provider's tables are called `table_types` now: an +`offset_provider_type=` keyword and an `.offset_provider_type` attribute are renamed, and so +are `OffsetProviderType` (now `TableTypes`) and `is_offset_provider_type`. What cannot be +rewritten from the source alone is reported instead: offset providers keyed by strings, which +become keyed by the connectivity (`{E2C: table}`), `.value` on a dimension (now `.tag`), and +`isinstance` checks against `Dimension`. + +The output is not formatted; run `ruff format` on the changed files afterwards. +""" + +from __future__ import annotations + +import ast +import dataclasses +import difflib +import pathlib +import re +from collections.abc import Iterable, Iterator + +import typer + + +cli = typer.Typer(no_args_is_help=True, name="migrate-connectivities", help=__doc__) + + +@dataclasses.dataclass(frozen=True) +class Edit: + """Replace source lines `[start, end)` (0-based) with `text`.""" + + start: int + end: int + text: str + + +@dataclasses.dataclass +class Module: + path: pathlib.Path + source: str + tree: ast.Module + edits: list[Edit] = dataclasses.field(default_factory=list) + notes: list[str] = dataclasses.field(default_factory=list) + #: Names of `gt4py.next` the migrated declarations use unqualified, to be imported. + needed: set[str] = dataclasses.field(default_factory=set) + #: Whether the migrated declarations need `typing` imported (for `Local: TypeAlias = ...`). + needed_typing: bool = False + + @property + def lines(self) -> list[str]: + return self.source.splitlines(keepends=True) + + def segment(self, node: ast.AST) -> str: + segment = ast.get_source_segment(self.source, node) + assert segment is not None + return segment + + def note(self, node: ast.stmt | ast.expr, message: str) -> None: + self.notes.append(f"{self.path}:{node.lineno}: {message}") + + +def _callee_name(call: ast.Call) -> tuple[str, str] | None: + """`(prefix, name)` of a call to `[prefix.]name`, e.g. `("gtx.", "Dimension")`.""" + match call.func: + case ast.Name(id=name): + return "", name + case ast.Attribute(value=value, attr=name): + return f"{ast.unparse(value)}.", name + return None + + +def _keyword(call: ast.Call, name: str, position: int) -> ast.expr | None: + for keyword in call.keywords: + if keyword.arg == name: + return keyword.value + return call.args[position] if len(call.args) > position else None + + +_DECLARATION_CALLEES = ("Dimension", "FieldOffset") + + +def _imported_aliases(module: Module) -> dict[str, str]: + """Local names of imported `Dimension` / `FieldOffset`, e.g. `{"FO": "FieldOffset"}`.""" + return { + alias.asname or alias.name: alias.name + for statement in ast.walk(module.tree) + if isinstance(statement, ast.ImportFrom) + for alias in statement.names + if alias.name in _DECLARATION_CALLEES + } + + +def _declarations(module: Module) -> Iterator[tuple[ast.Assign, str, ast.Call, str, str]]: + """Module-level `name = [prefix.]Dimension(...)` / `FieldOffset(...)` statements.""" + aliases = _imported_aliases(module) + for statement in module.tree.body: + if ( + isinstance(statement, ast.Assign) + and len(statement.targets) == 1 + and isinstance(statement.targets[0], ast.Name) + and isinstance(statement.value, ast.Call) + and (callee := _callee_name(statement.value)) is not None + ): + prefix, name = callee + name = aliases.get(name, name) if prefix == "" else name + if name in _DECLARATION_CALLEES: + yield statement, statement.targets[0].id, statement.value, prefix, name + + +def _replace(module: Module, statement: ast.stmt, text: str) -> None: + assert statement.end_lineno is not None + module.edits.append(Edit(statement.lineno - 1, statement.end_lineno, text)) + + +def _is_local_kind(module: Module, call: ast.Call) -> bool: + kind = _keyword(call, "kind", 1) + return kind is not None and module.segment(kind).split(".")[-1] == "LOCAL" + + +def _dimension_name(node: ast.expr) -> str | None: + """`KDim` for `KDim` and `dims.KDim`.""" + match node: + case ast.Name(id=name) | ast.Attribute(attr=name): + return name + return None + + +@dataclasses.dataclass +class _Evidence: + """Why a dimension looks like a Cartesian axis, and why like a mesh location.""" + + axis: list[str] = dataclasses.field(default_factory=list) + location: list[str] = dataclasses.field(default_factory=list) + + +def _classify_dimensions(modules: list[Module]) -> dict[str, _Evidence]: + """Collect, per non-local dimension name, the evidence for axis and for mesh location.""" + evidence: dict[str, _Evidence] = {} + tags: dict[str, str] = {} # the old `Dimension("K")` value, for `"_StaggeredK"` + for module in modules: + for _, name, call, _, kind_of_call in _declarations(module): + if kind_of_call == "Dimension" and not _is_local_kind(module, call): + found = evidence.setdefault(name, _Evidence()) + if (kind := _keyword(call, "kind", 1)) is not None and module.segment( + kind + ).endswith("VERTICAL"): + found.axis.append("kind=VERTICAL") + if call.args and isinstance(tag := call.args[0], ast.Constant): + tags[str(tag.value)] = name + + def add(node: ast.expr | None, role: str, reason: str) -> None: + if node is not None and (name := _dimension_name(node)) in evidence: + getattr(evidence[name], role).append(reason) + + for module in modules: + for _, name, call, _, kind_of_call in _declarations(module): + if kind_of_call != "FieldOffset": + continue + source, target = _keyword(call, "source", 1), _keyword(call, "target", 2) + if not isinstance(target, ast.Tuple) or source is None: + continue + if len(target.elts) == 1 and module.segment(target.elts[0]) == module.segment(source): + add(source, "axis", f"Cartesian offset '{name}'") + elif len(target.elts) == 2: + add(source, "location", f"neighbor offset '{name}'") + add(target.elts[0], "location", f"neighbor offset '{name}'") + for node in ast.walk(module.tree): + match node: + case ast.Call(func=func, args=[first, *_]) if _dimension_name(func) in ( + "as_offset", + "flip_staggered", + ): + add(first, "axis", f"'{_dimension_name(func)}'") + case ast.BinOp( + left=left, op=ast.Add() | ast.Sub(), right=ast.Constant(value=int() | float()) + ): + add(left, "axis", "index arithmetic") + case ast.Constant(value=str() as text) if text.startswith("_Staggered"): + if (name := tags.get(text.removeprefix("_Staggered"))) is not None: + evidence[name].axis.append(f"staggered counterpart '{text}'") + return evidence + + +def _migrate_declarations( + module: Module, cartesian: dict[str, str], evidence: dict[str, _Evidence] +) -> None: + for statement, name, call, prefix, kind_of_call in _declarations(module): + if kind_of_call == "Dimension": + kind = _keyword(call, "kind", 1) + kind_src = module.segment(kind) if kind is not None else None + if _is_local_kind(module, call): + base = "LocalDimensionIndex" + text = f"class {name}({prefix}{base}): ...\n" + else: + found = evidence.get(name, _Evidence()) + if found.axis and found.location: + base = "DimensionIndex" + module.note( + statement, + f"'{name}': conflicting evidence, decide by hand -- a Cartesian axis" + f" ({', '.join(found.axis)}) and a mesh location" + f" ({', '.join(found.location)}); declared as 'DimensionIndex'.", + ) + elif found.axis: + base = "CartesianAxisIndex" + else: + base = "DimensionIndex" + if not found.location: + module.note( + statement, + f"'{name}': declared as 'DimensionIndex'; declare it as" + " 'CartesianAxisIndex' if it is a Cartesian axis (index arithmetic," + " 'Staggered').", + ) + kind_arg = f", kind={kind_src}" if kind_src is not None else "" + text = f"class {name}({prefix}{base}{kind_arg}): ...\n" + if not prefix: + module.needed.add(base) + _replace(module, statement, text) + continue + + source, target = _keyword(call, "source", 1), _keyword(call, "target", 2) + if source is None or not isinstance(target, ast.Tuple): + module.note(statement, f"'{name}': unrecognized 'FieldOffset' arguments, not migrated.") + continue + if len(target.elts) == 2: + domain, local = (module.segment(element) for element in target.elts) + text = ( + f"class {name}({prefix}NeighborConnectivity[{domain}, {module.segment(source)}]):\n" + # NOTE: `TypeAlias`, not a plain assignment: it is what keeps the adopted local + # dimension a *type* for mypy (see ADR 0030). + f" Local: typing.TypeAlias = {local}\n" + ) + module.needed_typing = True + if not prefix: + module.needed.add("NeighborConnectivity") + _replace(module, statement, text) + elif len(target.elts) == 1 and module.segment(target.elts[0]) == module.segment(source): + # A Cartesian offset has no declaration any more: `Off[i]` is `Dim + i`. + cartesian[name] = module.segment(source) + _replace(module, statement, "") + else: + module.note(statement, f"'{name}': a cross-dimension offset has no class equivalent.") + + +def _is_local_kind_member(node: ast.expr) -> bool: + """`DimensionKind.LOCAL`, however `DimensionKind` is qualified.""" + return ( + isinstance(node, ast.Attribute) + and node.attr == "LOCAL" + and _dimension_name(node.value) == "DimensionKind" + ) + + +def _migrate_local_kind(module: Module) -> None: + """ + Drop module-level aliases of the removed `DimensionKind.LOCAL`; report its other uses. + + A local dimension is a `LocalDimensionIndex` subclass now, and its `kind` is `None`, so + `DimensionKind` has no `LOCAL` member left to name. + """ + in_declarations = { + id(node) for _, _, call, _, _ in _declarations(module) for node in ast.walk(call) + } + aliases = set() + for statement in module.tree.body: + if ( + isinstance(statement, ast.Assign) + and len(statement.targets) == 1 + and isinstance(statement.targets[0], ast.Name) + and _is_local_kind_member(statement.value) + ): + aliases.add(id(statement.value)) + _replace(module, statement, "") + module.note( + statement, + f"'{statement.targets[0].id}' removed: 'DimensionKind' has no 'LOCAL' member any" + " more; any other use of it has to test 'common.is_local_dimension(dim)'.", + ) + for node in ast.walk(module.tree): + if ( + _is_local_kind_member(node) + and id(node) not in in_declarations + and id(node) not in aliases + ): + module.note( + node, + "'DimensionKind.LOCAL' is removed: a local dimension is a 'LocalDimensionIndex'" + " subclass; test 'common.is_local_dimension(dim)' instead of its kind.", + ) + + +def _migrate_imports(module: Module) -> None: + """Drop imports of the removed `FieldOffset`; import the class names used unqualified.""" + fieldoffset_names = { + local for local, name in _imported_aliases(module).items() if name == "FieldOffset" + } + imported = { + alias.asname or alias.name + for statement in ast.walk(module.tree) + if isinstance(statement, ast.ImportFrom) + for alias in statement.names + } + missing = sorted(module.needed - imported) + typing_missing = module.needed_typing and not any( + isinstance(statement, ast.Import) and any(a.name == "typing" for a in statement.names) + for statement in ast.walk(module.tree) + ) + added = False + for statement in module.tree.body: + if not isinstance(statement, ast.ImportFrom): + continue + names = {alias.asname or alias.name for alias in statement.names} + drops = names & fieldoffset_names + adds_here = bool(missing) and not added and bool(names & {"Dimension", "DimensionKind"}) + if not (drops or adds_here): + continue + kept = [ + ast.unparse(alias) + for alias in statement.names + if (alias.asname or alias.name) not in fieldoffset_names + ] + text = ( + f"from {'.' * statement.level}{statement.module or ''} import {', '.join(kept)}\n" + if kept + else "" + ) + if adds_here: + text += f"from gt4py.next import {', '.join(missing)}\n" + added = True + _replace(module, statement, text) + if typing_missing: + # `Local: TypeAlias = ...` needs it; the first import statement is a safe place + for statement in module.tree.body: + if isinstance(statement, (ast.Import, ast.ImportFrom)): + module.edits.append( + Edit(statement.lineno - 1, statement.lineno - 1, "import typing\n") + ) + break + else: + module.notes.append( + f"{module.path}: import 'typing', used by the migrated declarations." + ) + if missing and not added: + module.notes.append( + f"{module.path}: import {', '.join(missing)} from 'gt4py.next', used by the migrated" + " declarations." + ) + + +class _CartesianUses(ast.NodeVisitor): + """Rewrite the uses of removed Cartesian offsets, and note the ones it cannot.""" + + def __init__(self, module: Module, cartesian: dict[str, str]) -> None: + self.module = module + self.cartesian = cartesian + #: (line, start column, end column, replacement) of single-line expression rewrites + self.rewrites: list[tuple[int, int, int, str]] = [] + self.handled: set[int] = set() + #: names bound in the enclosing function scopes, which shadow a removed offset + self.shadowed: list[set[str]] = [] + + def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef | ast.Lambda) -> None: + arguments = node.args + bound = { + argument.arg + for argument in (*arguments.posonlyargs, *arguments.args, *arguments.kwonlyargs) + } | { + name.id + for name in ast.walk(node) + if isinstance(name, ast.Name) and isinstance(name.ctx, ast.Store) + } + self.shadowed.append(bound) + self.generic_visit(node) + self.shadowed.pop() + + visit_AsyncFunctionDef = visit_FunctionDef + visit_Lambda = visit_FunctionDef + + def _rewrite(self, node: ast.expr, text: str) -> None: + assert node.end_lineno is not None and node.end_col_offset is not None + if node.lineno != node.end_lineno: + self.module.note(node, f"multi-line expression; rewrite by hand as '{text}'.") + return + self.rewrites.append((node.lineno - 1, node.col_offset, node.end_col_offset, text)) + + def _dimension_of(self, node: ast.expr) -> str | None: + """The dimension replacing `Off` or `module.Off`, qualified like the offset was.""" + match node: + case ast.Name(id=name) if name in self.cartesian and not any( + name in bound for bound in self.shadowed + ): + return self.cartesian[name] + case ast.Attribute(value=value, attr=name) if name in self.cartesian: + return f"{self.module.segment(value)}.{self.cartesian[name]}" + return None + + def visit_Subscript(self, node: ast.Subscript) -> None: + if (dim := self._dimension_of(node.value)) is not None: + match node.slice: + case ast.Constant(value=int() as index): + text = f"{dim} + {index}" if index >= 0 else f"{dim} - {-index}" + case ast.UnaryOp(op=ast.USub(), operand=ast.Constant(value=int() as index)): + text = f"{dim} - {index}" + case _: + text = f"{dim} + ({self.module.segment(node.slice)})" + self._rewrite(node, text) + self.handled.add(id(node.value)) + self.generic_visit(node) + + def visit_Name(self, node: ast.Name) -> None: + if id(node) not in self.handled and (dim := self._dimension_of(node)) is not None: + # e.g. `as_offset(Koff, field)`, which now takes the dimension + self._rewrite(node, dim) + + def visit_Attribute(self, node: ast.Attribute) -> None: + if id(node) not in self.handled and (dim := self._dimension_of(node)) is not None: + self._rewrite(node, dim) + else: + self.generic_visit(node) + + +def _migrate_cartesian_uses(module: Module, cartesian: dict[str, str]) -> None: + if not cartesian: + return + visitor = _CartesianUses(module, cartesian) + import_lines: set[int] = set() + for statement in ast.walk(module.tree): + if isinstance(statement, ast.ImportFrom) and any( + alias.name in cartesian for alias in statement.names + ): + import_lines.update(range(statement.lineno - 1, statement.end_lineno or 0)) + names = [ + cartesian.get(alias.name, alias.name) if alias.asname is None else alias.name + for alias in statement.names + ] + unique = list(dict.fromkeys(names)) + _replace( + module, + statement, + " " * statement.col_offset + + f"from {'.' * statement.level}{statement.module or ''} import {', '.join(unique)}\n", + ) + for statement in module.tree.body: + if not isinstance(statement, (ast.Import, ast.ImportFrom)): + visitor.visit(statement) + match statement: + case ast.Assign(targets=[ast.Name(id="__all__")], value=ast.List(elts=all_names)): + for entry in all_names: + if isinstance(entry, ast.Constant) and entry.value in cartesian: + module.note(entry, f"'__all__' lists the removed offset '{entry.value}'.") + + lines = module.lines + for line, start, end, text in sorted(visitor.rewrites, reverse=True): + if line in import_lines: + continue + lines[line] = lines[line][:start] + text + lines[line][end:] + module.source = "".join(lines) + + +#: Renamed names of `gt4py.next`, as keywords or attributes. +_RENAMED = { + "offset_provider_type": "table_types", + "OffsetProviderType": "TableTypes", + "is_offset_provider_type": "is_table_types", +} +#: Those also renamed where they are imported and used unqualified; `offset_provider_type` is not, +#: since a bare name of that spelling is as likely to be a variable of the user's own. +_RENAMED_IMPORTS = {"OffsetProviderType", "is_offset_provider_type"} + + +def _migrate_renamed_names(module: Module) -> None: + """ + Rename `offset_provider_type` and its relatives where they are keywords or attributes. + + A keyword is left alone, and noted, in a module that defines a function with a parameter of + that name: the call may be to that function rather than to `gt4py.next`. + """ + own_parameters = { + argument.arg + for node in ast.walk(module.tree) + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.Lambda)) + for argument in (*node.args.posonlyargs, *node.args.args, *node.args.kwonlyargs) + } + rewrites: list[tuple[int, int, int, str]] = [] + for node in ast.walk(module.tree): + match node: + case ast.keyword(arg=str() as name) if name in _RENAMED: + assert node.end_lineno is not None + if name in own_parameters: + module.note(node, f"'{name}=': renamed '{_RENAMED[name]}=' if it is gt4py's.") + continue + rewrites.append( + (node.lineno - 1, node.col_offset, node.col_offset + len(name), _RENAMED[name]) + ) + case ast.Attribute(attr=name) if name in _RENAMED: + assert node.end_lineno is not None and node.end_col_offset is not None + start = node.end_col_offset - len(name) + rewrites.append((node.end_lineno - 1, start, node.end_col_offset, _RENAMED[name])) + case ast.alias(name=name) if name in _RENAMED_IMPORTS and node.asname is None: + assert node.end_col_offset is not None + rewrites.append( + (node.lineno - 1, node.col_offset, node.end_col_offset, _RENAMED[name]) + ) + case ast.Name(id=name) if name in _RENAMED_IMPORTS: + assert node.end_col_offset is not None + rewrites.append( + (node.lineno - 1, node.col_offset, node.end_col_offset, _RENAMED[name]) + ) + lines = module.lines + for line, start, end, text in sorted(rewrites, reverse=True): + lines[line] = lines[line][:start] + text + lines[line][end:] + module.source = "".join(lines) + + +_PROVIDER_KEY_RE = re.compile(r"""(?P["'])(?P[A-Za-z_]\w*)(?P=quote)\s*:""") + + +def _report( + module: Module, + offset_keys: dict[str, str], + cartesian: dict[str, str], + dimension_names: set[str], +) -> None: + for number, line in enumerate(module.source.splitlines(), start=1): + for match in _PROVIDER_KEY_RE.finditer(line): + if (connectivity := offset_keys.get(match["name"])) is None: + continue + if connectivity in cartesian: + message = ( + f"offset-provider key '{match['name']}': remove the entry, a Cartesian shift" + f" ('{cartesian[connectivity]} + i') needs none." + ) + else: + message = ( + f"offset-provider key '{match['name']}' is keyed by the connectivity class" + f" now, e.g. '{{{connectivity}: table}}'." + ) + module.notes.append(f"{module.path}:{number}: {message}") + for node in ast.walk(module.tree): + match node: + case ast.Attribute( + value=ast.Name(id=name) | ast.Attribute(attr=name), attr="value" + ) if name in dimension_names: + module.note(node, f"'{name}.value': a dimension's name is '{name}.tag' now.") + case ast.Call( + func=ast.Name(id="isinstance"), args=[_, ast.Attribute(attr="Dimension")] + ): + module.note( + node, + "'isinstance(..., Dimension)': a dimension is a class now; use" + " 'isinstance(obj, gt4py.next.common.DimensionMeta)'.", + ) + + +def _apply(module: Module) -> str: + lines = module.lines + for edit in sorted(module.edits, key=lambda edit: edit.start, reverse=True): + lines[edit.start : edit.end] = [edit.text] if edit.text else [] + return "".join(lines) + + +def migrate(sources: dict[pathlib.Path, str]) -> tuple[dict[pathlib.Path, str], list[str]]: + """Migrate the given modules; return their new sources and what is left to do by hand.""" + modules = [ + Module(path=path, source=source, tree=ast.parse(source)) for path, source in sources.items() + ] + cartesian: dict[str, str] = {} + #: offset-provider keys that name a `FieldOffset`: its variable name and its tag, if different + offset_keys: dict[str, str] = {} + dimension_names: set[str] = set() + evidence = _classify_dimensions(modules) + for module in modules: + for _, name, call, _, kind_of_call in _declarations(module): + if kind_of_call == "Dimension": + dimension_names.add(name) + continue + offset_keys[name] = name + if call.args and isinstance(tag := call.args[0], ast.Constant): + offset_keys[str(tag.value)] = name + _migrate_declarations(module, cartesian, evidence) + _migrate_local_kind(module) + _migrate_imports(module) + + results: dict[pathlib.Path, str] = {} + notes: list[str] = [] + for module in modules: + _report(module, offset_keys, cartesian, dimension_names) + migrated = _apply(module) + # Cartesian uses are rewritten on the migrated text, re-parsed, so that line numbers + # refer to what the declaration edits left. + second = Module(path=module.path, source=migrated, tree=ast.parse(migrated)) + _migrate_cartesian_uses(second, cartesian) + rewritten = _apply(second) + third = Module(path=module.path, source=rewritten, tree=ast.parse(rewritten)) + _migrate_renamed_names(third) + results[module.path] = third.source + notes += module.notes + second.notes + third.notes + return results, notes + + +def _python_files(paths: Iterable[pathlib.Path]) -> Iterator[pathlib.Path]: + for path in paths: + if path.is_dir(): + yield from sorted(path.rglob("*.py")) + else: + yield path + + +@cli.command() +def run( + paths: list[pathlib.Path], + write: bool = typer.Option(False, "--write", help="Rewrite files instead of printing a diff."), +) -> None: + """Migrate the Python files in PATHS (directories are searched recursively).""" + sources = {path: path.read_text() for path in _python_files(paths)} + results, notes = migrate(sources) + for path, new in results.items(): + if new == sources[path]: + continue + if write: + path.write_text(new) + typer.echo(f"migrated {path}") + else: + typer.echo( + "".join( + difflib.unified_diff( + sources[path].splitlines(keepends=True), + new.splitlines(keepends=True), + fromfile=str(path), + tofile=str(path), + ) + ) + ) + if notes: + typer.echo("\nLeft to migrate by hand (line numbers of the original files):", err=True) + for note in notes: + typer.echo(f" {note}", err=True) + + +if __name__ == "__main__": + cli() diff --git a/scripts/tests/python/test_migrate_connectivities.py b/scripts/tests/python/test_migrate_connectivities.py new file mode 100644 index 0000000000..6845e1f6db --- /dev/null +++ b/scripts/tests/python/test_migrate_connectivities.py @@ -0,0 +1,267 @@ +# +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause +# + +"""Tests for the ``migrate_connectivities`` dev script.""" + +from __future__ import annotations + +import pathlib +import textwrap + +import migrate_connectivities + + +DIMENSIONS = textwrap.dedent( + """\ + import gt4py.next as gtx + + KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) + EdgeDim = gtx.Dimension("Edge") + CellDim = gtx.Dimension("Cell") + CEDim = gtx.Dimension("CE") + E2CDim = gtx.Dimension("E2C", gtx.DimensionKind.LOCAL) + C2EDim = gtx.Dimension("C2E", gtx.DimensionKind.LOCAL) + E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim, E2CDim)) + C2E = gtx.FieldOffset("C2E", source=EdgeDim, target=(CellDim, C2EDim)) + C2CE = gtx.FieldOffset("C2CE", source=CEDim, target=(CellDim, C2EDim)) + Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) + """ +) + +STENCIL = textwrap.dedent( + """\ + import gt4py.next as gtx + from gt4py.next.ffront.experimental import as_offset + + from pkg import dimension as dims + from pkg.dimension import E2C, KDim, Koff + + + def stencil(a, k): + b = a(Koff[1]) + a(Koff[-1]) + a(Koff[k]) + a(dims.Koff[1]) + return b(as_offset(Koff, k)) + b(as_offset(dims.Koff, k)) + + + def run(prog, grid): + prog(offset_provider={"E2C": grid.e2c, "Koff": KDim}) + return KDim.value + """ +) + + +def _migrate(**sources: str) -> tuple[dict[str, str], list[str]]: + results, notes = migrate_connectivities.migrate( + {pathlib.Path(name): source for name, source in sources.items()} + ) + return {str(path): source for path, source in results.items()}, notes + + +def test_declarations(): + results, _ = _migrate(dimension=DIMENSIONS) + + assert results["dimension"] == textwrap.dedent( + """\ + import typing + import gt4py.next as gtx + + class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... + class EdgeDim(gtx.DimensionIndex): ... + class CellDim(gtx.DimensionIndex): ... + class CEDim(gtx.DimensionIndex): ... + class E2CDim(gtx.LocalDimensionIndex): ... + class C2EDim(gtx.LocalDimensionIndex): ... + class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + Local: typing.TypeAlias = E2CDim + class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + Local: typing.TypeAlias = C2EDim + class C2CE(gtx.NeighborConnectivity[CellDim, CEDim]): + Local: typing.TypeAlias = C2EDim + """ + ) + + +def test_declarations_run(): + results, _ = _migrate(dimension=DIMENSIONS) + namespace: dict = {"__name__": "migrated_dimension"} + exec(results["dimension"], namespace) + + assert namespace["C2E"].Local is namespace["C2EDim"] + assert namespace["C2EDim"].owner is namespace["C2E"] + # `C2CE` shares `C2E`'s local dimension and is named by its own tag + assert namespace["C2CE"].Local is namespace["C2EDim"] + assert namespace["C2CE"].offset_tag == namespace["C2CE"].tag + + +def test_cartesian_offset_uses_across_modules(): + results, _ = _migrate(dimension=DIMENSIONS, stencil=STENCIL) + + stencil = results["stencil"] + assert "from pkg.dimension import E2C, KDim\n" in stencil + assert "a(KDim + 1) + a(KDim - 1) + a(KDim + (k)) + a(dims.KDim + 1)" in stencil + assert "b(as_offset(KDim, k)) + b(as_offset(dims.KDim, k))" in stencil + assert "Koff" not in stencil.replace('"Koff"', "") + + +def test_what_is_left_is_reported(): + _, notes = _migrate(dimension=DIMENSIONS, stencil=STENCIL) + + assert any("offset-provider key 'E2C'" in note for note in notes) + assert any("offset-provider key 'Koff'" in note for note in notes) + assert any("'KDim.value'" in note for note in notes) + + +BARE = textwrap.dedent( + """\ + from gt4py.next import Dimension, DimensionKind, FieldOffset as FO + + LOCAL = DimensionKind.LOCAL + Vertex = Dimension("Vertex") + Edge = Dimension("Edge") + V2EDim = Dimension("V2E", LOCAL) + V2E = FO("V2E_TAG", source=Edge, target=(Vertex, V2EDim)) + """ +) + + +def test_unqualified_names_and_aliases(): + results, notes = _migrate(bare=BARE, user='table = {"V2E_TAG": t}\n') + migrated = results["bare"] + + assert "FO" not in migrated + assert "from gt4py.next import Dimension, DimensionKind\n" in migrated + assert ( + "from gt4py.next import DimensionIndex, LocalDimensionIndex, NeighborConnectivity\n" + in migrated + ) + assert "class V2EDim(LocalDimensionIndex): ..." in migrated + assert "class V2E(NeighborConnectivity[Vertex, Edge]):" in migrated + assert " Local: typing.TypeAlias = V2EDim" in migrated + namespace: dict = {"__name__": "migrated_bare"} + exec(migrated, namespace) + assert namespace["V2E"].Local is namespace["V2EDim"] + # a key spelled with the offset's tag is reported, naming the class + assert any("'V2E_TAG'" in note and "{V2E: table}" in note for note in notes) + + +def test_shadowed_names_and_all_are_left_alone(): + source = textwrap.dedent( + """\ + from pkg.dimension import KDim, Koff + + __all__ = ["Koff"] + + + def helper(Koff): + return Koff + 1 + + + def uses(a): + return a(Koff[1]) + a(dims.KDim.value) + """ + ) + results, notes = _migrate(dimension=DIMENSIONS, user=source) + + assert "return Koff + 1" in results["user"] + assert "a(KDim + 1)" in results["user"] + assert any("'__all__' lists the removed offset 'Koff'" in note for note in notes) + assert any("'KDim.value'" in note for note in notes) + + +def test_table_types_rename(): + gt4py_calls = textwrap.dedent( + """\ + from gt4py.next import common + from gt4py.next.common import OffsetProviderType + + + def types_of(args) -> OffsetProviderType: + assert common.is_offset_provider_type(args.offset_provider_type) + return infer(program, offset_provider_type=args.offset_provider_type) + """ + ) + own_function = textwrap.dedent( + """\ + def run(program, offset_provider_type): ... + + + run(program, offset_provider_type=types) + """ + ) + results, notes = _migrate(gt4py_calls=gt4py_calls, own_function=own_function) + + assert results["gt4py_calls"] == textwrap.dedent( + """\ + from gt4py.next import common + from gt4py.next.common import TableTypes + + + def types_of(args) -> TableTypes: + assert common.is_table_types(args.table_types) + return infer(program, table_types=args.table_types) + """ + ) + # the keyword may be the user's own parameter: left alone, and noted + assert results["own_function"] == own_function + assert any("'offset_provider_type='" in note for note in notes) + + +AXES = textwrap.dedent( + """\ + import gt4py.next as gtx + from gt4py.next.ffront.experimental import as_offset + + IDim = gtx.Dimension("I") + JDim = gtx.Dimension("J") + HDim = gtx.Dimension("H") + XDim = gtx.Dimension("X") + Cell = gtx.Dimension("Cell") + Odd = gtx.Dimension("Odd") + Unknown = gtx.Dimension("Unknown") + C2ODim = gtx.Dimension("C2O", gtx.DimensionKind.LOCAL) + C2O = gtx.FieldOffset("C2O", source=Odd, target=(Cell, C2ODim)) + Ooff = gtx.FieldOffset("Ooff", source=Odd, target=(Odd,)) + + IHalf = gtx.Dimension("_StaggeredI") + + + def f(a, b, k): + return a(JDim + 1), b(as_offset(HDim, k)), gtx.flip_staggered(XDim) + """ +) + + +def test_axis_or_mesh_location_is_decided_per_declaration(): + results, notes = _migrate(axes=AXES) + migrated = results["axes"] + + # Cartesian axes: a staggered counterpart, index arithmetic, `as_offset`, `flip_staggered` + for axis in ("IDim", "JDim", "HDim", "XDim"): + assert f"class {axis}(gtx.CartesianAxisIndex): ..." in migrated + # a mesh location: the target dimension of a neighbor offset + assert "class Cell(gtx.DimensionIndex): ..." in migrated + # evidence for both is reported, not decided + assert "class Odd(gtx.DimensionIndex): ..." in migrated + assert any("'Odd': conflicting evidence" in note for note in notes) + # no evidence: a `DimensionIndex`, reported + assert "class Unknown(gtx.DimensionIndex): ..." in migrated + assert any( + "'Unknown': declared as 'DimensionIndex'; declare it as 'CartesianAxisIndex'" in note + for note in notes + ) + assert not any("'Cell'" in note or "'IDim'" in note for note in notes) + + +def test_local_kind_aliases_are_removed_and_other_uses_reported(): + source = BARE + "is_local = Vertex.kind == DimensionKind.LOCAL\n" + results, notes = _migrate(bare=source) + + assert "LOCAL = DimensionKind.LOCAL" not in results["bare"] + assert any("'LOCAL' removed" in note for note in notes) + assert any("bare:8: 'DimensionKind.LOCAL' is removed" in note for note in notes) diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index 253e574f95..00803ce442 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -52,7 +52,6 @@ from .ffront import fbuiltins from .ffront.decorator import field_operator, program, scan_operator from .ffront.fbuiltins import ( - FieldOffset, IndexType, abs, # noqa: A004 # shadowing arccos, @@ -153,7 +152,6 @@ "as_field", "as_connectivity", # from ffront - "FieldOffset", "field_operator", "program", "scan_operator", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 471b4a7abb..e3159b8b68 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -428,7 +428,6 @@ def staggered_base_tag(tag: Tag) -> Optional[Tag]: return match["base"] if (match := _STAGGERED_TAG_RE.match(tag)) is not None else None -@functools.cache def resolve(tag: Tag) -> Dimension: """ Return the dimension class a tag names, by importing it. @@ -441,6 +440,10 @@ def resolve(tag: Tag) -> Dimension: so the longest importable prefix wins and the remainder is walked as attributes. A collision would need a module path and an attribute chain to have the same spelling. + Only where the module path ends is memoized; the attribute walk is repeated on every call, so + a declaration redefined under the same name (e.g. by re-running a notebook cell) resolves to + the new class. + Parametrized dimensions such as `Staggered[K]` have no importable qualname; their tag has the form `[]` and is resolved by subscripting the owner, which goes through its intern table and so returns the identical class. @@ -466,27 +469,22 @@ def resolve(tag: Tag) -> Dimension: ) return owner[resolve(match["base"])] # type: ignore[index] # a StaggeredMeta, checked - parts = tag.split(".") - for split in range(len(parts), 0, -1): - try: - obj: Any = importlib.import_module(".".join(parts[:split])) - except ImportError: - continue - for attr in parts[split:]: - try: - obj = getattr(obj, attr) - except AttributeError as ex: - raise ValueError( - f"Cannot resolve dimension tag '{tag}': '{'.'.join(parts[:split])}' has" - f" no attribute '{attr}'." - ) from ex - if not isinstance(obj, DimensionMeta): - raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") - return cast(Dimension, obj) - raise ValueError( - f"Cannot resolve dimension tag '{tag}': no importable module prefix. A dimension" - " referenced from the IR must be declared at module level in an importable module." - ) + obj = _import_qualified_name(tag) + if not isinstance(obj, DimensionMeta): + raise ValueError(f"Tag '{tag}' resolves to '{obj}', which is not a dimension.") + return cast(Dimension, obj) + + +def _import_qualified_name_or_none(tag: Tag) -> Any: + """`_import_qualified_name`, returning `None` for a name that does not resolve.""" + if (split := _split_qualified_name_or_none(tag)) is None: + return None + module_name, attrs = split + obj: Any = sys.modules.get(module_name) or importlib.import_module(module_name) + for attr in attrs: + if (obj := getattr(obj, attr, None)) is None: + return None + return obj def resolve_loaded(tag: Tag) -> Optional[Dimension]: @@ -509,6 +507,52 @@ def resolve_loaded(tag: Tag) -> Optional[Dimension]: return None +def _import_qualified_name(tag: Tag) -> Any: + """Import the object a dotted qualified name refers to; see `resolve`.""" + module_name, attrs = _split_qualified_name(tag) + obj: Any = sys.modules.get(module_name) or importlib.import_module(module_name) + for attr in attrs: + try: + obj = getattr(obj, attr) + except AttributeError as ex: + raise ValueError( + f"Cannot resolve tag '{tag}': '{module_name}' has no attribute '{'.'.join(attrs)}'." + ) from ex + return obj + + +@functools.cache +def _split_qualified_name_or_none(tag: Tag) -> Optional[tuple[str, tuple[str, ...]]]: + """ + `_split_qualified_name`, returning `None` instead of raising. + + Separate and memoized so that a string that is not a qualified name -- an offset-provider key + of hand-written IR, say -- costs one import attempt in total, not one per call. A module whose + import *fails* other than by not being found is reported, not cached away. + """ + try: + return _split_qualified_name(tag) + except ValueError: + return None + + +@functools.cache +def _split_qualified_name(tag: Tag) -> tuple[str, tuple[str, ...]]: + """Split a dotted name at its longest importable module prefix: `(module, attributes)`.""" + parts = tag.split(".") + for split in range(len(parts), 0, -1): + module_name = ".".join(parts[:split]) + try: + importlib.import_module(module_name) + except ImportError: + continue + return module_name, tuple(parts[split:]) + raise ValueError( + f"Cannot resolve tag '{tag}': no importable module prefix. A dimension or connectivity" + " referenced from the IR must be declared at module level in an importable module." + ) + + class Infinity(enum.Enum): """Describes an unbounded `UnitRange`.""" @@ -1132,9 +1176,7 @@ def asnumpy(self) -> np.ndarray: ... def as_scalar(self) -> core_defs.ScalarT: ... @abc.abstractmethod - def premap( - self, index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity] - ) -> Field: ... + def premap(self, index_field: Connectivity | type[NeighborConnectivity]) -> Field: ... @abc.abstractmethod def restrict(self, item: AnyIndexSpec) -> Self: ... @@ -1142,8 +1184,8 @@ def restrict(self, item: AnyIndexSpec) -> Self: ... @abc.abstractmethod def __call__( self, - index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], - *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + index_field: Connectivity | type[NeighborConnectivity], + *args: Connectivity | type[NeighborConnectivity], ) -> Field: ... @abc.abstractmethod @@ -1543,10 +1585,18 @@ def is_neighbor_table(obj: Any) -> TypeGuard[NeighborTable]: OffsetProviderElem: TypeAlias = NeighborTable # Note: `OffsetProvider` and `TableTypes` should not be accessed directly, # use the `get_offset` and `get_offset_type` functions instead. +#: Neighbor tables keyed by the connectivity's `offset_tag`, which is how the IR names it. OffsetProvider: TypeAlias = Mapping[Tag, OffsetProviderElem] #: The types of an offset provider's tables, under the same keys: what transformations and code #: generation see instead of the tables (ADR 0019). TableTypes: TypeAlias = Mapping[Tag, NeighborTableType] +#: An offset provider as users write it: keyed by `NeighborConnectivity` declarations (or, at the +#: IR level, by tags). The entry points of a program normalize it to an `OffsetProvider` with +#: `as_tag_keyed_offset_provider`, so everything below them sees tags only. +#: NOTE: `Any` keys, since `Mapping` is invariant in its key type: a tag-keyed provider would not +#: be a `Mapping[type[NeighborConnectivity] | Tag, ...]`. Keys are checked at runtime instead. +OffsetProviderLike: TypeAlias = Mapping[Any, OffsetProviderElem] +TableTypesLike: TypeAlias = Mapping[Any, NeighborTableType] def is_offset_provider(obj: Any) -> TypeGuard[OffsetProvider]: @@ -1606,8 +1656,12 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid """ # TODO(havogt): Once we have a custom class for `OffsetProvider`, we can absorb this functionality into it. if offset_tag not in offset_provider: - raise KeyError(f"Offset '{offset_tag}' not found in offset provider.") - return offset_provider[offset_tag] # TODO return a valid dimension + raise KeyError( + f"Connectivity '{offset_tag}' not found in the offset provider, which has" + f" {sorted(map(str, offset_provider))}. Offset providers are keyed by" + " 'NeighborConnectivity' declarations, e.g. '{V2E: v2e_table}'." + ) + return offset_provider[offset_tag] get_offset_type: Callable[[TableTypes, str], NeighborTableType] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and TableTypes overlap @@ -1664,7 +1718,9 @@ def has_offset(offset_provider: OffsetProvider | TableTypes, offset_tag: str) -> return True -def hash_offset_provider_items_by_id(offset_provider: OffsetProvider) -> int: +def hash_offset_provider_items_by_id( + offset_provider: OffsetProviderLike | TableTypesLike, +) -> int: """ Compute hash of an offset provider on the tuples of key and value id. @@ -1755,8 +1811,8 @@ def inverse_image(self, image_range: UnitRange | NamedRange) -> Sequence[NamedRa def premap( self, - index_field: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], - *args: Connectivity | fbuiltins.FieldOffset | type[NeighborConnectivity], + index_field: Connectivity | type[NeighborConnectivity], + *args: Connectivity | type[NeighborConnectivity], ) -> Connectivity: raise NotImplementedError() @@ -2207,7 +2263,7 @@ 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)] + return cls.bound_table()[cls._local()(int(item))] if "Local" in cls.__dict__: raise TypeError( f"'{cls.__qualname__}[{item!r}]': a connectivity is indexed by an integer" @@ -2224,20 +2280,32 @@ 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`. + The type of the connectivity in DSL code: a shift from `Codomain` to `(Domain, Local)`. + + Its tag is `offset_tag`, which is how the IR names the connectivity and how the offset + provider is keyed once normalized (see `as_tag_keyed_offset_provider`). """ - from gt4py.next.ffront import fbuiltins + from gt4py.next.type_system import type_specifications as ts + + local = cls._local() + return ts.ShiftType(codomain=cls.codomain, domain=(cls.domain, local), tag=cls.offset_tag) - if (field_offset := cls.__dict__.get("_field_offset")) is None: - field_offset = fbuiltins.FieldOffset( - cls.offset_tag, source=cls.codomain, target=(cls.domain, cls._local()) + def bound_table(cls) -> NeighborTable: + """The neighbor table bound to this connectivity in the current embedded execution.""" + from gt4py.next import embedded + + offset_provider = embedded.context.get_offset_provider(None) + if offset_provider is None: + raise RuntimeError( + f"'{cls.__qualname__}' can only be resolved to a table during embedded execution." + ) + table = get_offset(offset_provider, cls.offset_tag) + if not is_neighbor_table(table): + raise TypeError( + f"'{cls.__qualname__}' is bound to '{table}', which is not a neighbor table." ) - type.__setattr__(cls, "_field_offset", field_offset) - return field_offset + return table class NeighborConnectivity[Domain: DimensionIndex, Codomain: DimensionIndex]( @@ -2416,18 +2484,35 @@ def fail(reason: str) -> NoReturn: table_type = _unbound_table_type(table) else: fail(f"expected a neighbor table, got '{table}'") - if isinstance(table_type.connectivity, ConnectivityMeta) and ( - table_type.connectivity is not connectivity + + def redefined(found: Sequence[Any], expected: Sequence[Any], what: str = "dimension") -> str: + if any(f is not e and f.tag == e.tag for f, e in zip(found, expected)): + return ( + f" (a {what} of the same name but a different class: was the declaration" + " redefined, e.g. by re-running a notebook cell?)" + ) + return "" + + if isinstance(bound_to := table_type.connectivity, ConnectivityMeta) and ( + bound_to is not connectivity ): - fail(f"its type is bound to '{table_type.connectivity.__qualname__}'") + fail( + f"its type is bound to '{bound_to.__qualname__}'" + + redefined((bound_to,), (connectivity,), what="connectivity") + ) + expected_domain = (connectivity.domain, local) if tuple(table_type.domain) != expected_domain: fail( f"its domain is '({', '.join(map(str, table_type.domain))})'," f" expected '({', '.join(map(str, expected_domain))})'" + + redefined(table_type.domain, expected_domain) ) if table_type.codomain is not connectivity.codomain: - fail(f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'") + fail( + f"its codomain is '{table_type.codomain}', expected '{connectivity.codomain}'" + + redefined((table_type.codomain,), (connectivity.codomain,)) + ) if not np.issubdtype(table_type.dtype.scalar_type, np.integer): fail(f"its dtype '{table_type.dtype}' is not integral") if local.max_neighbors is not None and table_type.max_neighbors != local.max_neighbors: @@ -2453,3 +2538,156 @@ def fail(reason: str) -> NoReturn: f" its neighbors, but the table has skip value {table_type.skip_value}" ) return dataclasses.replace(table_type, connectivity=connectivity) + + +@overload +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderLike, *, strict: bool = True +) -> OffsetProvider: ... +@overload +def as_tag_keyed_offset_provider( + offset_provider: TableTypesLike, *, strict: bool = True +) -> TableTypes: ... +def as_tag_keyed_offset_provider( + offset_provider: OffsetProviderLike | TableTypesLike, *, strict: bool = True +) -> OffsetProvider | TableTypes: + """ + Key an offset provider by tags, the form the IR and the backends use. + + A `NeighborConnectivity` key becomes its `offset_tag`: its local dimension's tag, or its own + for a connectivity sharing another one's local dimension. A string key is taken to be + such a tag already, and is rejected if it cannot be one: a tag is a qualified name, so a bare + name such as `"V2E"` is the removed `FieldOffset` spelling. + + Called on every program call, so it does not check tables against their declarations; see + `check_offset_provider`. + + Args: + offset_provider: The provider to normalize. + strict: Whether to reject string keys that cannot be tags. Internal hooks that are handed + hand-written providers, such as `embedded.context.update`, pass `False`. + """ + if not any(isinstance(key, ConnectivityMeta) for key in offset_provider): + if strict: + _check_tag_keys(offset_provider) + return offset_provider + result: dict[Tag, Any] = {} + for key, value in offset_provider.items(): + tag = key.offset_tag if isinstance(key, ConnectivityMeta) else key + if tag in result: + raise ValueError(f"The offset provider binds '{tag}' twice.") + result[tag] = value + if strict: + _check_tag_keys(result) + return result + + +def _check_tag_keys(offset_provider: Mapping[Any, Any]) -> None: + for key in offset_provider: + if not isinstance(key, str) or "." not in key: + raise TypeError( + f"Invalid offset-provider key {key!r}: offset providers are keyed by" + " 'NeighborConnectivity' declarations, e.g. '{V2E: v2e_table}'. A bare name is the" + " spelling of the removed 'FieldOffset' (see ADR 0030)." + ) + + +#: Offset providers already checked, by the hash of their `(key, id(table))` items. Bounded, and +#: not authoritative: like the compiled-program cache (which keys on the same hash), it can in +#: principle skip a check when a freed table is replaced at the same address. See +#: `check_offset_provider`. +_CHECKED_OFFSET_PROVIDERS: Final[collections.OrderedDict[int, None]] = collections.OrderedDict() +_CHECKED_OFFSET_PROVIDERS_MAX: Final = 256 + + +def check_offset_provider( + offset_provider: OffsetProviderLike | TableTypesLike, *, deep: bool = False +) -> None: + """ + Check every table of an offset provider against its connectivity declaration. + + A key that does not name a declared connectivity -- e.g. a tag used only by hand-written IR -- + has no declaration to be checked against and is skipped. Providers are remembered by the + identity of their tables, so repeated calls with the same tables cost one hash. + + Args: + offset_provider: The provider, keyed by declarations or by tags. + deep: Also compare the skip-value positions of tables over one shared local dimension, + which reads the tables. The compile path does; the call path does not. + + Raises: + ValueError: If a table does not match its declaration, see `check_neighbor_table`. + """ + if (seen := hash_offset_provider_items_by_id(offset_provider)) in _CHECKED_OFFSET_PROVIDERS: + return + for key, table in offset_provider.items(): + declaration: Any = key + if isinstance(key, str): + if (declaration := _import_qualified_name_or_none(key)) is None: + continue + if isinstance(declaration, DimensionMeta): + # the local dimension's tag names its owner's table + declaration = getattr(declaration, "owner", None) + if isinstance(declaration, ConnectivityMeta): + if declaration.offset_tag not in (key, getattr(key, "offset_tag", None)): + # e.g. `{V2E.tag: table}`: the connectivity's own tag, which is the IR name only + # of a connectivity *sharing* a local dimension + raise ValueError( + f"Invalid offset-provider key '{key}': it names the connectivity" + f" '{declaration.__qualname__}', whose key is the declaration itself" + f" ('{{{declaration.__qualname__}: table}}')." + ) + check_neighbor_table(cast(type[NeighborConnectivity], declaration), table) + _check_shared_local_dimensions(offset_provider, deep=deep) + _CHECKED_OFFSET_PROVIDERS[seen] = None + while len(_CHECKED_OFFSET_PROVIDERS) > _CHECKED_OFFSET_PROVIDERS_MAX: + _CHECKED_OFFSET_PROVIDERS.popitem(last=False) + + +def _check_shared_local_dimensions( + offset_provider: OffsetProviderLike | TableTypesLike, *, deep: bool = False +) -> None: + """ + Check that the tables over one local dimension have the same neighbor structure. + + Reductions and sparse fields take the neighbor count and the skip values of a local dimension + from any one table over it (see `connectivity_key_over`), which is only sound if all of them + agree: the same number of neighbors, and a skip value at the same positions. + """ + by_local_dim: dict[Tag, list[tuple[Any, Any]]] = collections.defaultdict(list) + for key, table in offset_provider.items(): + if (neighbor_dim := _neighbor_dim_of(table)) is not None: + by_local_dim[neighbor_dim.tag].append((key, table)) + for local_tag, tables in by_local_dim.items(): + (first_key, first), *others = tables + first_type = first if isinstance(first, NeighborTableType) else _unbound_table_type(first) + for key, table in others: + table_type = ( + table if isinstance(table, NeighborTableType) else _unbound_table_type(table) + ) + same_structure = (table_type.max_neighbors, table_type.has_skip_values) == ( + first_type.max_neighbors, + first_type.has_skip_values, + ) + if ( + deep + and same_structure + and first_type.has_skip_values + and is_neighbor_table(first) + and is_neighbor_table(table) + ): + # NOTE: compared where the tables live, without copying device arrays to the host. + xp = first.array_ns # type: ignore[attr-defined] # all tables are NdArrayFields + same_structure = first.ndarray.shape == table.ndarray.shape and bool( + xp.all( + (first.ndarray == first_type.skip_value) + == (xp.asarray(table.ndarray) == table_type.skip_value) + ) + ) + if not same_structure: + raise ValueError( + f"'{key}' and '{first_key}' are bound to tables over the same local dimension" + f" '{local_tag}' with a different neighbor structure: connectivities sharing a" + " local dimension must have the same number of neighbors, and skip values at" + " the same positions." + ) diff --git a/src/gt4py/next/embedded/context.py b/src/gt4py/next/embedded/context.py index 8183a3292c..6ca2544ce9 100644 --- a/src/gt4py/next/embedded/context.py +++ b/src/gt4py/next/embedded/context.py @@ -73,7 +73,7 @@ def get_offset_provider(default: _T = _NO_DEFAULT_SENTINEL) -> common.OffsetProv def update( *, closure_column_range: common.NamedRange | eve.NothingType = eve.NOTHING, - offset_provider: common.OffsetProvider | eve.NothingType = eve.NOTHING, + offset_provider: common.OffsetProviderLike | eve.NothingType = eve.NOTHING, ) -> Generator[None, None, None]: """Context handler updating the current embedded context with the provided values.""" @@ -83,7 +83,9 @@ def update( closure_token = gtx_embedded.context._closure_column_range.set(closure_column_range) if offset_provider is not eve.NOTHING: assert not isinstance(offset_provider, eve.NothingType) - offset_provider_token = gtx_embedded.context._offset_provider.set(offset_provider) + offset_provider_token = gtx_embedded.context._offset_provider.set( + common.as_tag_keyed_offset_provider(offset_provider, strict=False) + ) try: yield None diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index 26504c332d..24497b1c01 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -230,9 +230,7 @@ def as_scalar(self) -> core_defs.ScalarT: def premap( self: NdArrayField, - *connectivities: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], + *connectivities: common.Connectivity | type[common.NeighborConnectivity], ) -> NdArrayField: """ Rearrange the field content using the provided connectivities (index mappings). @@ -307,13 +305,9 @@ def premap( codomains_counter: collections.Counter[common.Dimension] = collections.Counter() for connectivity in connectivities: - # For neighbor reductions, a FieldOffset or a connectivity declaration is passed - # instead of an actual Connectivity + # For neighbor reductions, a connectivity declaration is passed instead of a table 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() + connectivity = connectivity.bound_table() assert isinstance(connectivity, common.Connectivity) # Current implementation relies on skip_value == -1: @@ -362,10 +356,8 @@ def premap( def __call__( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: return functools.reduce( lambda field, current_index_field: field.premap(current_index_field), @@ -939,15 +931,12 @@ def _concat_where( NdArrayField.register_builtin_func(experimental.concat_where, _concat_where) # type: ignore[arg-type] -def _as_offset(offset: fbuiltins.FieldOffset, offset_field: NdArrayField) -> common.Connectivity: - if not fbuiltins.is_cartesian_offset(offset): - target_dims = ", ".join(d.__qualname__ for d in offset.target) # for the diagnostic - raise ValueError( - f"'as_offset' is only supported for Cartesian offsets " - f"(single target dimension equal to source dimension); " - f"got source '{offset.source.__qualname__}' and target ({target_dims})." - ) - source_dim = offset.source +def _as_offset(source_dim: common.Dimension, offset_field: NdArrayField) -> common.Connectivity: + if not ( + isinstance(source_dim, common.DimensionMeta) + and issubclass(source_dim, common.AnyCartesianAxisIndex) + ): + raise TypeError(f"'as_offset' shifts along a Cartesian axis, got '{source_dim}'.") coords = _identity_index_array( offset_field.domain, source_dim, offset_field.array_ns, dtype=fbuiltins.IndexType ) diff --git a/src/gt4py/next/ffront/decorator.py b/src/gt4py/next/ffront/decorator.py index a464de9763..c5e9f1534b 100644 --- a/src/gt4py/next/ffront/decorator.py +++ b/src/gt4py/next/ffront/decorator.py @@ -159,9 +159,9 @@ def _make_compiled_programs_pool( def compile( self, - offset_provider: common.TableTypes - | common.OffsetProvider - | list[common.TableTypes | common.OffsetProvider] + offset_provider: common.TableTypesLike + | common.OffsetProviderLike + | list[common.TableTypesLike | common.OffsetProviderLike] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -195,7 +195,8 @@ def compile( self.compilation_options.connectivities if offset_provider is None else offset_provider ) if not isinstance(offset_provider, list): - offset_provider = [offset_provider] # type: ignore[list-item] # cleanup offset_provider vs offset_provider_type + offset_provider = [offset_provider] # type: ignore[list-item] # cleanup offset_provider vs table_types + offset_provider = [common.as_tag_keyed_offset_provider(op) for op in offset_provider] assert all( common.is_offset_provider(op) or common.is_table_types(op) for op in offset_provider @@ -372,12 +373,13 @@ def with_bound_args(self, **kwargs: Any) -> ProgramWithBoundArgs: def __call__( self, *args: Any, - offset_provider: common.OffsetProvider | None = None, + offset_provider: common.OffsetProviderLike | None = None, enable_jit: bool | None = None, **kwargs: Any, ) -> None: - if offset_provider is None: - offset_provider = {} + offset_provider = ( + {} if offset_provider is None else common.as_tag_keyed_offset_provider(offset_provider) + ) enable_jit = self.compilation_options.enable_jit if enable_jit is None else enable_jit with program_call_context( @@ -408,6 +410,7 @@ def __call__( stacklevel=2, ) + common.check_offset_provider(offset_provider) with next_embedded.context.update(offset_provider=offset_provider): with embedded_program_call_context(self, args, offset_provider, kwargs): self.definition_stage.definition(*args, **kwargs) @@ -427,7 +430,7 @@ class ProgramWithBoundArgs(Program): @override def __call__( - self, *args: Any, offset_provider: common.OffsetProvider | None = None, **kwargs: Any + self, *args: Any, offset_provider: common.OffsetProviderLike | None = None, **kwargs: Any ) -> None: if offset_provider is None: offset_provider = {} @@ -484,9 +487,9 @@ def __call__( @override def compile( self, - offset_provider: common.TableTypes - | common.OffsetProvider - | list[common.TableTypes | common.OffsetProvider] + offset_provider: common.TableTypesLike + | common.OffsetProviderLike + | list[common.TableTypesLike | common.OffsetProviderLike] | None = None, **static_args: list[xtyping.MaybeNestedInTuple[core_defs.Scalar]], ) -> Self: @@ -651,7 +654,9 @@ def __gt_closure_vars__(self) -> dict[str, Any]: def __call__(self, *args: Any, enable_jit: bool | None = None, **kwargs: Any) -> Any: if not next_embedded.context.within_valid_context() and self.backend is not None: # non embedded execution - offset_provider = {**kwargs.pop("offset_provider", {})} + offset_provider = { + **common.as_tag_keyed_offset_provider(kwargs.pop("offset_provider", {})) + } if "out" not in kwargs: raise errors.MissingArgumentError(None, "out", True) out = kwargs.pop("out") @@ -673,7 +678,10 @@ def __call__(self, *args: Any, enable_jit: bool | None = None, **kwargs: Any) -> else: if not next_embedded.context.within_valid_context(): # field_operator as program - kwargs["offset_provider"] = {**kwargs.pop("offset_provider", {})} + kwargs["offset_provider"] = { + **common.as_tag_keyed_offset_provider(kwargs.pop("offset_provider", {})) + } + common.check_offset_provider(kwargs["offset_provider"]) attributes = ( self.definition_stage.attributes if self.definition_stage @@ -717,6 +725,11 @@ class FieldOperatorFromFoast(FieldOperator): @override def __call__(self, *args: Any, **kwargs: Any) -> Any: assert self.backend is not None + if "offset_provider" in kwargs: + kwargs["offset_provider"] = common.as_tag_keyed_offset_provider( + kwargs["offset_provider"] + ) + common.check_offset_provider(kwargs["offset_provider"]) compiled_fo = self.backend.compile( self.foast_stage, arguments.CompileTimeArgs.from_concrete(*args, **kwargs) ) diff --git a/src/gt4py/next/ffront/experimental.py b/src/gt4py/next/ffront/experimental.py index b547d67663..0ac7ebe4cf 100644 --- a/src/gt4py/next/ffront/experimental.py +++ b/src/gt4py/next/ffront/experimental.py @@ -10,11 +10,19 @@ from gt4py._core import definitions as core_defs from gt4py.next import common, named_collections -from gt4py.next.ffront.fbuiltins import BuiltInFunction, FieldOffset, WhereBuiltinFunction +from gt4py.next.ffront.fbuiltins import BuiltInFunction, WhereBuiltinFunction @BuiltInFunction -def as_offset(offset: FieldOffset, field: common.Field, /) -> common.Connectivity: +def as_offset( + dim: type[common.AnyCartesianAxisIndex], field: common.Field, / +) -> common.Connectivity: + """ + Shift along the Cartesian axis `dim` by the per-point amounts in the integer `field`. + + `a(as_offset(KDim, k_offsets))` reads `a` at `k + k_offsets[k]` in `KDim`. Like `KDim + 1`, + it needs index arithmetic, so `dim` must be a Cartesian axis (or its staggered partner). + """ raise NotImplementedError() diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index dc8552425c..f08475cf33 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -7,7 +7,6 @@ # SPDX-License-Identifier: BSD-3-Clause import dataclasses -import functools import inspect import math import operator @@ -35,7 +34,6 @@ from gt4py._core import definitions as core_defs from gt4py.next import common, named_collections from gt4py.next.common import Dimension, Field # noqa: F401 [unused-import] for TYPE_BUILTINS -from gt4py.next.iterator import runtime from gt4py.next.type_system import type_specifications as ts @@ -127,10 +125,14 @@ def _type_conversion_helper(t: type) -> type[ts.TypeSpec] | tuple[type[ts.TypeSpec], ...]: if t is common.Field: return ts.FieldType - elif t is common.Dimension: + elif t is common.Dimension or ( + # e.g. `type[AnyCartesianAxisIndex]`: a dimension narrowed to a level of the hierarchy, + # whose restriction the builtin's own type deduction checks + get_origin(t) is type + and isinstance(arg := get_args(t)[0], type) + and issubclass(arg, common.DimensionIndex) + ): return ts.DimensionType - elif t is FieldOffset: - return ts.ShiftType elif t is common.Connectivity: return ts.ShiftType elif t is core_defs.ScalarT: @@ -472,80 +474,3 @@ def impl( assert (diff := actual_export - should_export) == set(), ( f"Symbol(s) exported but not defined in 'fbuiltins': {diff}" ) - - -# TODO(tehrengruber): FieldOffset and runtime.Offset are not an exact conceptual -# match. Revisit if we want to continue subclassing here. If we split -# them also check whether Dimension should continue to be the shared or define -# guidelines for decision. -@dataclasses.dataclass(frozen=True) -class FieldOffset(runtime.Offset): - #: The tag, i.e. the offset-provider key. - value: str - source: common.Dimension - target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] - - @functools.cached_property - def _cache(self) -> dict: - return {} - - def __post_init__(self) -> None: - if len(self.target) == 2 and not common.is_local_dimension(self.target[1]): - raise ValueError("Second dimension in offset must be a local dimension.") - - def __gt_type__(self) -> ts.ShiftType: - return ts.ShiftType(codomain=self.source, domain=self.target, tag=self.value) - - @property - def Local(self) -> common.Dimension: - """The local dimension, as `V2E.Local` names it on a `NeighborConnectivity`.""" - if len(self.target) != 2: - raise AttributeError( - f"'{self.value}' is a Cartesian offset and has no local dimension." - ) - return self.target[1] - - def __getitem__(self, offset: int) -> common.Connectivity: - """Serve as a connectivity factory.""" - from gt4py.next import embedded # avoid circular import - - assert isinstance(self.value, str) - current_offset_provider = embedded.context.get_offset_provider(None) - assert current_offset_provider is not None - offset_definition = common.get_offset(current_offset_provider, self.value) - - assert common.is_neighbor_table(offset_definition) - named_index = (self.target[-1])(offset) - connectivity = offset_definition[named_index] - - return connectivity - - def as_connectivity_field(self) -> common.Connectivity: - """Convert to connectivity field using the offset providers in current embedded execution context.""" - from gt4py.next import embedded # avoid circular import - - assert isinstance(self.value, str) - current_offset_provider = embedded.context.get_offset_provider(None) - assert current_offset_provider is not None - offset_definition = common.get_offset(current_offset_provider, self.value) - - cache_key = id(offset_definition) - if (connectivity := self._cache.get(cache_key, None)) is None: - if isinstance(offset_definition, common.Connectivity): - connectivity = offset_definition - else: - raise NotImplementedError() - - self._cache[cache_key] = connectivity - - return connectivity - - -def is_cartesian_offset(offset: FieldOffset | ts.ShiftType) -> bool: - shift_type = offset.__gt_type__() if isinstance(offset, FieldOffset) else offset - return ( - len(shift_type.domain) == 1 - and shift_type.codomain == shift_type.domain[0] - and shift_type.codomain.kind == shift_type.domain[0].kind - and not common.is_local_dimension(shift_type.domain[0]) - ) diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index 1ef443c17c..e81ab4795c 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -809,20 +809,10 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call: ): raise errors.DSLError(node.location, "Functions can only be called directly.") elif isinstance(new_func.type, ts.FieldType): - for arg in new_args: - # A Cartesian `FieldOffset` shifts by the index it is subscripted with, so it - # carries no displacement on its own. Only an offset with a local dimension is - # meaningful unsubscripted, as the neighbor access `field(Off)`. - if ( - isinstance(arg, (foast.Name, foast.Attribute)) - and isinstance(arg.type, ts.ShiftType) - and len(arg.type.domain) == 1 - ): - raise errors.DSLError( - arg.location, - f"Cannot shift by the Cartesian offset '{arg!s}' without an index.", - hints=[f"Give the displacement, e.g. '{arg!s}[1]'."], - ) + # NOTE: no unsubscripted Cartesian offset to reject here: a Cartesian shift is + # `a(Dim + i)`, and the only other shift with a one-dimensional domain is the result + # of `as_offset`, which is a call, not a name. + pass elif isinstance(new_func.type, ts.DimensionType): assert common.is_local_dimension(new_func.type.dim) return foast.Call( @@ -1013,16 +1003,16 @@ def _visit_astype(self, node: foast.Call, **kwargs: Any) -> foast.Call: def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: arg_0 = node.args[0].type arg_1 = node.args[1].type - assert isinstance(arg_0, ts.ShiftType) assert isinstance(arg_1, ts.FieldType) - if not fbuiltins.is_cartesian_offset(arg_0): - domain_dims = ", ".join(d.__qualname__ for d in arg_0.domain) # for the diagnostic + if not isinstance(arg_0, ts.DimensionType) or not issubclass( + arg_0.dim, common.AnyCartesianAxisIndex + ): raise errors.DSLError( node.location, - f"'as_offset' is only supported for Cartesian offsets " - f"(a single domain dimension equal to the codomain); " - f"got codomain '{arg_0.codomain.__qualname__}' and domain ({domain_dims}).", + f"'as_offset' shifts along a Cartesian axis, e.g. 'as_offset(KDim, field)';" + f" got '{arg_0}'.", ) + dim = arg_0.dim if not type_info.is_integral(arg_1): raise errors.DSLError( node.location, @@ -1031,16 +1021,20 @@ def _visit_as_offset(self, node: foast.Call, **kwargs: Any) -> foast.Call: f"{node.location}", ) - if arg_0.codomain not in arg_1.dims: + if dim not in arg_1.dims: raise errors.DSLError( node.location, f"Incompatible argument in call to '{node.func!s}': " - f"'{arg_0.codomain}' not in list of offset field dimensions '{arg_1.dims}'. " + f"'{dim}' not in list of offset field dimensions '{arg_1.dims}'. " f"{node.location}", ) return foast.Call( - func=node.func, args=node.args, kwargs=node.kwargs, type=arg_0, location=node.location + func=node.func, + args=node.args, + kwargs=node.kwargs, + type=ts.ShiftType(codomain=dim, domain=(dim,)), + location=node.location, ) def _deduce_where_return_type( diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index aa6731e11c..4e32756098 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -330,9 +330,9 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: # `field(as_offset(Off, offset_field))` case foast.Call(func=foast.Name(id="as_offset")): func_args = arg - offset_type = func_args.args[0].type - assert isinstance(offset_type, ts.ShiftType) - dim = offset_type.codomain + dim_type = func_args.args[0].type + assert isinstance(dim_type, ts.DimensionType) + dim = dim_type.dim offset_field = self.visit(func_args.args[1], **kwargs) current_expr = im.as_fieldop( im.lambda_("__it", "__offset")( diff --git a/src/gt4py/next/ffront/past_to_itir.py b/src/gt4py/next/ffront/past_to_itir.py index 595ac45afd..7572f5ec15 100644 --- a/src/gt4py/next/ffront/past_to_itir.py +++ b/src/gt4py/next/ffront/past_to_itir.py @@ -17,7 +17,6 @@ from gt4py.eve import NodeTranslator, traits from gt4py.next import common, config, errors, utils from gt4py.next.ffront import ( - fbuiltins, gtcallable, program_ast as past, stages as ffront_stages, @@ -74,7 +73,7 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> stages.CompilableProgramDef: """ all_closure_vars = transform_utils._get_closure_vars_recursively(inp.data.closure_vars) offsets_and_dimensions = transform_utils._filter_closure_vars_by_type( - all_closure_vars, fbuiltins.FieldOffset, common.ConnectivityMeta, common.DimensionMeta + all_closure_vars, 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 2b441796a6..e01e07d6c7 100644 --- a/src/gt4py/next/ffront/transform_utils.py +++ b/src/gt4py/next/ffront/transform_utils.py @@ -10,7 +10,6 @@ from typing import Any, Iterable, Optional from gt4py.next import common -from gt4py.next.ffront import fbuiltins from gt4py.next.ffront.gtcallable import GTCallable @@ -47,9 +46,7 @@ def _filter_closure_vars_by_type(closure_vars: dict[str, Any], *types: type) -> def _deduce_grid_type( requested_grid_type: Optional[common.GridType], - offsets_and_dimensions: Iterable[ - fbuiltins.FieldOffset | type[common.NeighborConnectivity] | common.Dimension - ], + offsets_and_dimensions: Iterable[type[common.NeighborConnectivity] | common.Dimension], ) -> common.GridType: """ Derive grid type from actually occurring dimensions and check against optional user request. @@ -61,9 +58,7 @@ def _deduce_grid_type( deduced_grid_type = common.GridType.CARTESIAN for o in offsets_and_dimensions: - if isinstance(o, common.ConnectivityMeta) or ( - isinstance(o, fbuiltins.FieldOffset) and not fbuiltins.is_cartesian_offset(o) - ): + if isinstance(o, common.ConnectivityMeta): deduced_grid_type = common.GridType.UNSTRUCTURED break if isinstance(o, common.DimensionMeta) and common.is_local_dimension(o): @@ -75,7 +70,7 @@ def _deduce_grid_type( and deduced_grid_type == common.GridType.UNSTRUCTURED ): raise ValueError( - "'grid_type == GridType.CARTESIAN' was requested, but unstructured 'FieldOffset' or local 'Dimension' was found." + "'grid_type == GridType.CARTESIAN' was requested, but a 'NeighborConnectivity' or a local dimension was found." ) return deduced_grid_type if requested_grid_type is None else requested_grid_type diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 73c5e10102..d3078fd224 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -48,7 +48,6 @@ exceptions as embedded_exceptions, operators, ) -from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import builtins, runtime from gt4py.next.type_system import type_specifications as ts, type_translation @@ -144,9 +143,7 @@ def asnumpy(self) -> np.ndarray: def premap( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError @@ -164,10 +161,8 @@ def as_scalar(self) -> typing.Never: def __call__( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: raise NotImplementedError() @@ -1147,10 +1142,8 @@ def as_scalar(self) -> core_defs.IntegralScalar: def premap( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1290,10 +1283,8 @@ def asnumpy(self) -> np.ndarray: def premap( self, - index_field: common.Connectivity - | fbuiltins.FieldOffset - | type[common.NeighborConnectivity], - *args: common.Connectivity | fbuiltins.FieldOffset | type[common.NeighborConnectivity], + index_field: common.Connectivity | type[common.NeighborConnectivity], + *args: common.Connectivity | type[common.NeighborConnectivity], ) -> common.Field: # TODO can be implemented by constructing and ndarray (but do we know of which kind?) raise NotImplementedError() @@ -1388,10 +1379,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 @@ -1446,9 +1437,22 @@ def __gt_type__(self) -> ts.ListType: ) +def _as_offset_tag( + offset: runtime.Offset | type[common.NeighborConnectivity] | OffsetPart, +) -> OffsetPart: + if isinstance(offset, common.ConnectivityMeta): + return offset.offset_tag + return offset.value if isinstance(offset, runtime.Offset) else offset + + @builtins.neighbors.register(EMBEDDED) -def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: - offset_str = offset.value if isinstance(offset, runtime.Offset) else offset +def neighbors(offset: runtime.Offset | type[common.NeighborConnectivity], it: ItIterator) -> _List: + field_offset: runtime.Offset = ( + runtime.Offset(value=offset.offset_tag) + if isinstance(offset, common.ConnectivityMeta) + else offset + ) + offset_str = _as_offset_tag(field_offset) assert isinstance(offset_str, str) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None @@ -1460,7 +1464,7 @@ def neighbors(offset: runtime.Offset, it: ItIterator) -> _List: for i in range(len(connectivity.domain[1].unit_range)) if (shifted := it.shift(offset_str, i)).can_deref() ), - offset=offset, + offset=field_offset, ) @@ -1676,9 +1680,9 @@ def _dimension_to_tag( return {k.tag: v for k, v in domain.items()} -def _validate_domain(domain: Domain, offset_provider_type: common.TableTypes) -> None: +def _validate_domain(domain: Domain, table_types: common.TableTypes) -> None: if isinstance(domain, runtime.CartesianDomain): - if any(isinstance(o, common.NeighborTableType) for o in offset_provider_type.values()): + if any(isinstance(o, common.NeighborTableType) for o in table_types.values()): raise RuntimeError( "Got a 'CartesianDomain', but found a 'Connectivity' in 'offset_provider', expected 'UnstructuredDomain'." ) diff --git a/src/gt4py/next/iterator/runtime.py b/src/gt4py/next/iterator/runtime.py index 3b604fcb31..05b66fac2f 100644 --- a/src/gt4py/next/iterator/runtime.py +++ b/src/gt4py/next/iterator/runtime.py @@ -78,7 +78,11 @@ def __call__( offset_provider=None, column_axis=None, ): - offset_provider = offset_provider or self.offset_provider + # NOTE: not strict: iterator IR names offsets by arbitrary strings. + offset_provider = common.as_tag_keyed_offset_provider( + offset_provider or self.offset_provider, strict=False + ) + common.check_offset_provider(offset_provider) column_axis = column_axis or self.column_axis if backend is not None: @@ -133,7 +137,7 @@ def fendef( ) -def _deduce_domain(domain: dict[common.Dimension, range], offset_provider_type: common.TableTypes): +def _deduce_domain(domain: dict[common.Dimension, range], table_types: common.TableTypes): if isinstance(domain, UnstructuredDomain): domain_builtin = builtins.unstructured_domain elif isinstance(domain, CartesianDomain): @@ -141,7 +145,7 @@ def _deduce_domain(domain: dict[common.Dimension, range], offset_provider_type: else: domain_builtin = ( builtins.unstructured_domain - if any(isinstance(o, common.NeighborTableType) for o in offset_provider_type.values()) + if any(isinstance(o, common.NeighborTableType) for o in table_types.values()) else builtins.cartesian_domain ) diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index a6ae605bab..fb4eeddb37 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=o.offset_tag) if callable(o): if o.__name__ == "": return lambdadef(o) diff --git a/src/gt4py/next/iterator/transforms/collapse_tuple.py b/src/gt4py/next/iterator/transforms/collapse_tuple.py index 1d049897d9..ebaf5ae85b 100644 --- a/src/gt4py/next/iterator/transforms/collapse_tuple.py +++ b/src/gt4py/next/iterator/transforms/collapse_tuple.py @@ -187,7 +187,7 @@ def apply( node: itir.Node, *, remove_letified_make_tuple_elements: bool = True, - offset_provider_type: Optional[common.TableTypes] = None, + table_types: Optional[common.TableTypes] = None, within_stencil: Optional[bool] = None, # manually passing enabled transformations is mostly for allowing separate testing of the modes enabled_transformations: Optional[Transformation] = None, @@ -211,7 +211,7 @@ def apply( point, without recursing into its children. """ enabled_transformations = enabled_transformations or cls.enabled_transformations - offset_provider_type = offset_provider_type or {} + table_types = table_types or {} if isinstance(node, itir.Program): within_stencil = False @@ -230,7 +230,7 @@ def apply( if requires_types: node = itir_type_inference.infer( node, - offset_provider_type=offset_provider_type, + table_types=table_types, allow_undeclared_symbols=allow_undeclared_symbols, ) diff --git a/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py b/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py index aaf34d5759..1561172531 100644 --- a/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py +++ b/src/gt4py/next/iterator/transforms/concat_where/expand_tuple_args.py @@ -27,12 +27,12 @@ def apply( cls, node: itir.Node, *, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, allow_undeclared_symbols: bool = False, ) -> itir.Node: node = type_inference.infer( node, - offset_provider_type=offset_provider_type, + table_types=table_types, allow_undeclared_symbols=allow_undeclared_symbols, ) return cls().visit(node) diff --git a/src/gt4py/next/iterator/transforms/cse.py b/src/gt4py/next/iterator/transforms/cse.py index 78d4612a16..eb48753759 100644 --- a/src/gt4py/next/iterator/transforms/cse.py +++ b/src/gt4py/next/iterator/transforms/cse.py @@ -462,7 +462,7 @@ def apply( cls, node: ProgramOrExpr, within_stencil: bool | None = None, - offset_provider_type: common.TableTypes | None = None, + table_types: common.TableTypes | None = None, *, uids: utils.IDGeneratorPool, ) -> ProgramOrExpr: @@ -475,9 +475,9 @@ def apply( "The expression's context must be specified using `within_stencil`." ) - offset_provider_type = offset_provider_type or {} + table_types = table_types or {} node = itir_type_inference.infer( - node, offset_provider_type=offset_provider_type, allow_undeclared_symbols=not is_program + node, table_types=table_types, allow_undeclared_symbols=not is_program ) return cls(uids=uids).visit(node, within_stencil=within_stencil) diff --git a/src/gt4py/next/iterator/transforms/dead_code_elimination.py b/src/gt4py/next/iterator/transforms/dead_code_elimination.py index 8a7d84b0f2..eb1e5d83b5 100644 --- a/src/gt4py/next/iterator/transforms/dead_code_elimination.py +++ b/src/gt4py/next/iterator/transforms/dead_code_elimination.py @@ -17,7 +17,7 @@ def dead_code_elimination( program: itir.Program, *, uids: utils.IDGeneratorPool, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, ) -> itir.Program: """ Perform dead code elimination on a program by simplifying or removing @@ -74,7 +74,7 @@ def dead_code_elimination( program, enabled_transformations=~CollapseTuple.Transformation.PROPAGATE_TO_IF_ON_TUPLES, uids=uids, - offset_provider_type=offset_provider_type, + table_types=table_types, ) # type: ignore[assignment] # always an itir.Program return program diff --git a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py index 7028e0837e..4d1aed0e42 100644 --- a/src/gt4py/next/iterator/transforms/expand_tuple_maps.py +++ b/src/gt4py/next/iterator/transforms/expand_tuple_maps.py @@ -61,12 +61,12 @@ def apply( node: ProgramOrExpr, *, uids: utils.IDGeneratorPool | None, - offset_provider_type: common.TableTypes | None = None, + table_types: common.TableTypes | None = None, ) -> ProgramOrExpr: if node.type is None: node = itir_inference.infer( node, - offset_provider_type=offset_provider_type or {}, + table_types=table_types or {}, allow_undeclared_symbols=not isinstance(node, itir.Program), ) if uids is None: diff --git a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py index 9babe30fe8..0c25c2d601 100644 --- a/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py +++ b/src/gt4py/next/iterator/transforms/fuse_as_fieldop.py @@ -123,7 +123,7 @@ def fuse_as_fieldop( expr: itir.Expr, eligible_args: list[bool], *, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, enable_cse: bool, uids: utils.IDGeneratorPool, ) -> itir.Expr: @@ -192,7 +192,7 @@ def fuse_as_fieldop( if enable_cse: # TODO(havogt): We should investigate how to keep the tree small without having to run CSE. new_node = cse.CommonSubexpressionElimination.apply( - new_node, within_stencil=False, uids=uids, offset_provider_type=offset_provider_type + new_node, within_stencil=False, uids=uids, table_types=table_types ) return new_node @@ -273,7 +273,7 @@ class FuseAsFieldOp( >>> print( ... FuseAsFieldOp.apply( ... nested_as_fieldop, - ... offset_provider_type={}, + ... table_types={}, ... allow_undeclared_symbols=True, ... uids=utils.IDGeneratorPool(), ... ) @@ -301,7 +301,7 @@ def all(self) -> FuseAsFieldOp.Transformation: enabled_transformations = Transformation.all() uids: utils.IDGeneratorPool - offset_provider_type: common.TableTypes + table_types: common.TableTypes enable_cse: bool # option to disable is mainly for testing purposes @classmethod @@ -309,7 +309,7 @@ def apply( cls, node: itir.Program, *, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, uids: utils.IDGeneratorPool, allow_undeclared_symbols=False, within_set_at_expr: Optional[bool] = None, @@ -320,7 +320,7 @@ def apply( node = type_inference.infer( node, - offset_provider_type=offset_provider_type, + table_types=table_types, allow_undeclared_symbols=allow_undeclared_symbols, ) @@ -330,7 +330,7 @@ def apply( new_node = cls( uids=uids, enabled_transformations=enabled_transformations, - offset_provider_type=offset_provider_type, + table_types=table_types, enable_cse=enable_cse, ).visit(node, within_set_at_expr=within_set_at_expr) # The `FuseAsFieldOp` pass does not fully preserve the type information yet. In particular @@ -338,7 +338,7 @@ def apply( # everything here ensuring later passes can use the information. new_node = type_inference.infer( new_node, - offset_provider_type=offset_provider_type, + table_types=table_types, allow_undeclared_symbols=allow_undeclared_symbols, ) return new_node @@ -421,7 +421,7 @@ def transform_fuse_as_fieldop(self, node: itir.Node, **kwargs): node, eligible_els, uids=self.uids, - offset_provider_type=self.offset_provider_type, + table_types=self.table_types, enable_cse=self.enable_cse, ), **{**kwargs, "recurse": False}, diff --git a/src/gt4py/next/iterator/transforms/global_tmps.py b/src/gt4py/next/iterator/transforms/global_tmps.py index 8952554479..75d0bd447e 100644 --- a/src/gt4py/next/iterator/transforms/global_tmps.py +++ b/src/gt4py/next/iterator/transforms/global_tmps.py @@ -336,7 +336,7 @@ def create_global_tmps( keep_existing_domains=True, ) program = type_inference.infer( - program, offset_provider_type=common.offset_provider_to_type(offset_provider) + program, table_types=common.offset_provider_to_type(offset_provider) ) declarations = program.declarations.copy() diff --git a/src/gt4py/next/iterator/transforms/infer_domain.py b/src/gt4py/next/iterator/transforms/infer_domain.py index 8df09b0860..97c3575dca 100644 --- a/src/gt4py/next/iterator/transforms/infer_domain.py +++ b/src/gt4py/next/iterator/transforms/infer_domain.py @@ -486,9 +486,7 @@ def infer_expr( if not revisit_already_inferred and hasattr(expr.annex, "domain"): return expr, {} - itir_type_inference.reinfer( - expr, offset_provider_type=common.offset_provider_to_type(offset_provider) - ) + itir_type_inference.reinfer(expr, table_types=common.offset_provider_to_type(offset_provider)) el_types, domain = gtx_utils.equalize_tuple_structure( gtx_utils.tree_map( collection_type=ts.TupleType, result_collection_constructor=lambda _, elts: tuple(elts) @@ -589,7 +587,7 @@ def infer_program( ) program = itir_type_inference.infer( - program, offset_provider_type=common.offset_provider_to_type(offset_provider) + program, table_types=common.offset_provider_to_type(offset_provider) ) return itir.Program( diff --git a/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py b/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py index 489d140027..4cf15cb182 100644 --- a/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py +++ b/src/gt4py/next/iterator/transforms/inline_dynamic_shifts.py @@ -33,17 +33,17 @@ def _dynamic_shift_args(node: itir.Expr) -> list[bool] | None: @dataclasses.dataclass class InlineDynamicShifts(eve.NodeTranslator, eve.VisitorWithSymbolTableTrait): - offset_provider_type: common.TableTypes + table_types: common.TableTypes uids: utils.IDGeneratorPool @classmethod def apply( cls, node: itir.Program, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, uids: utils.IDGeneratorPool, ): - return cls(offset_provider_type=offset_provider_type, uids=uids).visit(node) + return cls(table_types=table_types, uids=uids).visit(node) def visit_FunCall(self, node: itir.FunCall, **kwargs): node = self.generic_visit(node, **kwargs) @@ -83,7 +83,7 @@ def visit_FunCall(self, node: itir.FunCall, **kwargs): expr, fuse_args, uids=self.uids, - offset_provider_type=self.offset_provider_type, + table_types=self.table_types, enable_cse=True, ) diff --git a/src/gt4py/next/iterator/transforms/inline_scalar.py b/src/gt4py/next/iterator/transforms/inline_scalar.py index 223a484702..530fdd7abb 100644 --- a/src/gt4py/next/iterator/transforms/inline_scalar.py +++ b/src/gt4py/next/iterator/transforms/inline_scalar.py @@ -19,8 +19,8 @@ class InlineScalar(eve.NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) @classmethod - def apply(cls, program: itir.Program, offset_provider_type: common.TableTypes): - program = itir_inference.infer(program, offset_provider_type=offset_provider_type) + def apply(cls, program: itir.Program, table_types: common.TableTypes): + program = itir_inference.infer(program, table_types=table_types) return cls().visit(program) def generic_visit(self, node, **kwargs): diff --git a/src/gt4py/next/iterator/transforms/pass_manager.py b/src/gt4py/next/iterator/transforms/pass_manager.py index 0865ffda00..408edf10f8 100644 --- a/src/gt4py/next/iterator/transforms/pass_manager.py +++ b/src/gt4py/next/iterator/transforms/pass_manager.py @@ -151,7 +151,7 @@ def apply_common_transforms( # relies on static information or `symbolic_domain_sizes`. assert common.is_offset_provider(offset_provider) - offset_provider_type = common.offset_provider_to_type(offset_provider) + table_types = common.offset_provider_to_type(offset_provider) symbolic_domain_sizes = _process_symbolic_domains_option( ir, offset_provider, symbolic_domain_sizes, use_max_domain_range_on_unstructured_shift @@ -169,15 +169,13 @@ def apply_common_transforms( # test_can_deref. We didn't notice previously as FieldOpFusion did this implicitly everywhere. ir = inline_lifts.InlineLifts().visit(ir) - ir = concat_where.expand_tuple_args(ir, offset_provider_type=offset_provider_type) # type: ignore[assignment] # always an itir.Program - ir = expand_tuple_maps.ExpandTupleMaps.apply( - ir, uids=uids, offset_provider_type=offset_provider_type - ) + ir = concat_where.expand_tuple_args(ir, table_types=table_types) # type: ignore[assignment] # always an itir.Program + ir = expand_tuple_maps.ExpandTupleMaps.apply(ir, uids=uids, table_types=table_types) ir = dead_code_elimination.dead_code_elimination( - ir, uids=uids, offset_provider_type=offset_provider_type + ir, uids=uids, table_types=table_types ) # domain inference does not support dead-code ir = inline_dynamic_shifts.InlineDynamicShifts.apply( - ir, offset_provider_type=offset_provider_type, uids=uids + ir, table_types=table_types, uids=uids ) # domain inference does not support dynamic offsets yet ir = infer_domain_ops.InferDomainOps.apply(ir) ir = concat_where.canonicalize_domain_argument(ir) @@ -205,17 +203,15 @@ def apply_common_transforms( inlined, enabled_transformations=~CollapseTuple.Transformation.PROPAGATE_TO_IF_ON_TUPLES, uids=uids, - offset_provider_type=offset_provider_type, + table_types=table_types, ) # type: ignore[assignment] # always an itir.Program - inlined = InlineScalar.apply(inlined, offset_provider_type=offset_provider_type) + inlined = InlineScalar.apply(inlined, table_types=table_types) # This pass is required to run after CollapseTuple as otherwise we can not inline # expressions like `tuple_get(make_tuple(as_fieldop(stencil)(...)))` where stencil returns # a list. Such expressions must be inlined however because no backend supports such # field operators right now. - inlined = fuse_as_fieldop.FuseAsFieldOp.apply( - inlined, uids=uids, offset_provider_type=offset_provider_type - ) + inlined = fuse_as_fieldop.FuseAsFieldOp.apply(inlined, uids=uids, table_types=table_types) if inlined == ir: break @@ -225,14 +221,12 @@ def apply_common_transforms( # breaks in test_zero_dim_tuple_arg as trivial tuple_get is not inlined if common_subexpression_elimination: - ir = CommonSubexpressionElimination.apply( - ir, offset_provider_type=offset_provider_type, uids=uids - ) + ir = CommonSubexpressionElimination.apply(ir, table_types=table_types, uids=uids) ir = MergeLet().visit(ir) ir = InlineLambdas.apply(ir, opcount_preserving=True) if extract_temporaries: - ir = infer(ir, inplace=True, offset_provider_type=offset_provider_type) + ir = infer(ir, inplace=True, table_types=table_types) ir = global_tmps.create_global_tmps( ir, offset_provider=offset_provider, @@ -247,7 +241,7 @@ def apply_common_transforms( if unroll_reduce: for _ in range(10): - unrolled = UnrollReduce.apply(ir, offset_provider_type=offset_provider_type, uids=uids) + unrolled = UnrollReduce.apply(ir, table_types=table_types, uids=uids) unrolled = CollapseListGet().visit(unrolled) unrolled = NormalizeShifts().visit(unrolled) # this is required as nested neighbor reductions can contain lifts, e.g., @@ -276,7 +270,7 @@ def apply_fieldview_transforms( # to work with / translate domains. use_max_domain_range_on_unstructured_shift: Optional[bool] = None, ) -> itir.Program: - offset_provider_type = common.offset_provider_to_type(offset_provider) + table_types = common.offset_provider_to_type(offset_provider) uids = utils.IDGeneratorPool() @@ -287,16 +281,12 @@ def apply_fieldview_transforms( ir = inline_fundefs.InlineFundefs().visit(ir) ir = inline_fundefs.prune_unreferenced_fundefs(ir) # required for dead-code-elimination and `prune_empty_concat_where` pass - ir = concat_where.expand_tuple_args(ir, offset_provider_type=offset_provider_type) # type: ignore[assignment] # always an itir.Program - ir = expand_tuple_maps.ExpandTupleMaps.apply( - ir, uids=uids, offset_provider_type=offset_provider_type - ) + ir = concat_where.expand_tuple_args(ir, table_types=table_types) # type: ignore[assignment] # always an itir.Program + ir = expand_tuple_maps.ExpandTupleMaps.apply(ir, uids=uids, table_types=table_types) - ir = dead_code_elimination.dead_code_elimination( - ir, offset_provider_type=offset_provider_type, uids=uids - ) + ir = dead_code_elimination.dead_code_elimination(ir, table_types=table_types, uids=uids) ir = inline_dynamic_shifts.InlineDynamicShifts.apply( - ir, offset_provider_type=offset_provider_type, uids=uids + ir, table_types=table_types, uids=uids ) # domain inference does not support dynamic offsets yet ir = infer_domain_ops.InferDomainOps.apply(ir) diff --git a/src/gt4py/next/iterator/transforms/unroll_reduce.py b/src/gt4py/next/iterator/transforms/unroll_reduce.py index 79067bf663..03fed83779 100644 --- a/src/gt4py/next/iterator/transforms/unroll_reduce.py +++ b/src/gt4py/next/iterator/transforms/unroll_reduce.py @@ -52,7 +52,7 @@ def _get_partial_local_dims(reduce_args: Iterable[itir.Expr]) -> Iterable[common def _get_connectivity( applied_reduce_node: itir.FunCall, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, ) -> common.NeighborTableType: """Return single connectivity that is compatible with the arguments of the reduce.""" if not cpm.is_applied_reduce(applied_reduce_node): @@ -61,7 +61,7 @@ def _get_connectivity( connectivities: list[common.NeighborTableType] = [] for local_dim in _get_partial_local_dims(applied_reduce_node.args): conn = common.get_offset_type( - offset_provider_type, common.connectivity_key_over(offset_provider_type, local_dim) + table_types, common.connectivity_key_over(table_types, local_dim) ) assert isinstance(conn, common.NeighborTableType) connectivities.append(conn) @@ -87,15 +87,13 @@ class UnrollReduce(PreserveLocationVisitor, NodeTranslator): def apply( cls, node: itir.Node, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, uids: utils.IDGeneratorPool, ) -> itir.Node: - return cls(uids=uids).visit(node, offset_provider_type=offset_provider_type) + return cls(uids=uids).visit(node, table_types=table_types) - def _visit_reduce( - self, node: itir.FunCall, offset_provider_type: common.TableTypes - ) -> itir.Expr: - connectivity_type = _get_connectivity(node, offset_provider_type) + def _visit_reduce(self, node: itir.FunCall, table_types: common.TableTypes) -> itir.Expr: + connectivity_type = _get_connectivity(node, table_types) max_neighbors = connectivity_type.max_neighbors has_skip_values = connectivity_type.has_skip_values diff --git a/src/gt4py/next/iterator/type_system/inference.py b/src/gt4py/next/iterator/type_system/inference.py index 77d0b033a8..4ad4e2e2a2 100644 --- a/src/gt4py/next/iterator/type_system/inference.py +++ b/src/gt4py/next/iterator/type_system/inference.py @@ -112,7 +112,7 @@ class ObservableTypeSynthesizer(type_synthesizer.TypeSynthesizer): >>> square_func_type_synthesizer = type_synthesizer.type_synthesizer( ... lambda base: power(base, int_type) ... ) - >>> square_func_type_synthesizer(float_type, offset_provider_type={}) + >>> square_func_type_synthesizer(float_type, table_types={}) ScalarType(kind=, shape=None) Note that without a corresponding call the function itself can not be fully typed and as such @@ -126,7 +126,7 @@ class ObservableTypeSynthesizer(type_synthesizer.TypeSynthesizer): ... node=square_func, ... store_inferred_type_in_node=True, ... ) - >>> o_type_synthesizer(float_type, offset_provider_type={}) + >>> o_type_synthesizer(float_type, table_types={}) ScalarType(kind=, shape=None) >>> square_func.type == ts.FunctionType( ... pos_only_args=[float_type], pos_or_kw_args={}, kw_only_args={}, returns=float_type @@ -182,16 +182,14 @@ def on_type_ready(self, cb: Callable[[ts.TypeSpec], None]) -> None: def __call__( self, *args: type_synthesizer.TypeOrTypeSynthesizer, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, **kwargs, ) -> Union[ts.TypeSpec, ObservableTypeSynthesizer]: assert all(isinstance(arg, (ts.TypeSpec, ObservableTypeSynthesizer)) for arg in args), ( "ObservableTypeSynthesizer can only be used with arguments that are TypeSpec or ObservableTypeSynthesizer" ) - return_type_or_synthesizer = self.type_synthesizer( - *args, offset_provider_type=offset_provider_type, **kwargs - ) + return_type_or_synthesizer = self.type_synthesizer(*args, table_types=table_types, **kwargs) # return type is a typing rule by itself if isinstance(return_type_or_synthesizer, type_synthesizer.TypeSynthesizer): @@ -256,7 +254,7 @@ class ITIRTypeInference(eve.NodeTranslator): PRESERVED_ANNEX_ATTRS = ("domain",) - offset_provider_type: Optional[common.TableTypes] + table_types: Optional[common.TableTypes] #: Allow sym refs to symbols that have not been declared. Mostly used in testing. allow_undeclared_symbols: bool #: Reinference-mode skipping already typed nodes. @@ -267,7 +265,7 @@ def apply( cls, node: T, *, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, inplace: bool = False, allow_undeclared_symbols: bool = False, ) -> T: @@ -278,7 +276,7 @@ def apply( node: The :class:`itir.Node` to infer the types of. Keyword Arguments: - offset_provider_type: Offset provider dictionary. + table_types: Offset provider dictionary. inplace: Write types directly to the given ``node`` instead of returning a copy. allow_undeclared_symbols: Allow references to symbols that don't have a corresponding declaration. This is useful for testing or inference on partially inferred sub-nodes. @@ -341,7 +339,7 @@ def apply( ) instance = cls( - offset_provider_type=offset_provider_type, + table_types=table_types, allow_undeclared_symbols=allow_undeclared_symbols, reinfer=False, ) @@ -351,9 +349,7 @@ def apply( return node @classmethod - def apply_reinfer( - cls, node: T, *, offset_provider_type: Optional[common.TableTypes] = None - ) -> T: + def apply_reinfer(cls, node: T, *, table_types: Optional[common.TableTypes] = None) -> T: """ Given a partially typed node infer the type of ``node`` and its sub-nodes. @@ -365,14 +361,12 @@ def apply_reinfer( Arguments: node: The :class:`itir.Node` to infer the types of. - offset_provider_type: Offset provider dictionary. + table_types: Offset provider dictionary. """ if node.type: # already inferred return node - instance = cls( - offset_provider_type=offset_provider_type, allow_undeclared_symbols=True, reinfer=True - ) + instance = cls(table_types=table_types, allow_undeclared_symbols=True, reinfer=True) instance.visit(node, ctx=_INITIAL_CONTEXT) return node @@ -563,7 +557,7 @@ def visit_FunCall( fun = self.visit(node.fun, ctx=ctx) args = self.visit(node.args, ctx=ctx) - result = fun(*args, **syntactic_info, offset_provider_type=self.offset_provider_type) + result = fun(*args, **syntactic_info, table_types=self.table_types) if isinstance(result, ObservableTypeSynthesizer): assert not result.node diff --git a/src/gt4py/next/iterator/type_system/type_synthesizer.py b/src/gt4py/next/iterator/type_system/type_synthesizer.py index 910dd68dcb..9a4785e90b 100644 --- a/src/gt4py/next/iterator/type_system/type_synthesizer.py +++ b/src/gt4py/next/iterator/type_system/type_synthesizer.py @@ -44,16 +44,16 @@ class TypeSynthesizer: - isinstance checks to determine if an object is actually (meant to be) a type synthesizer and not just any callable. - writing simple type synthesizers without cluttering the signature with the additional - offset_provider_type argument that is only needed by some. + table_types argument that is only needed by some. """ type_synthesizer: Callable[..., TypeOrTypeSynthesizer] cache: bool = False def __post_init__(self): - if "offset_provider_type" not in inspect.signature(self.type_synthesizer).parameters: + if "table_types" not in inspect.signature(self.type_synthesizer).parameters: synthesizer = self.type_synthesizer - self.type_synthesizer = lambda *args, offset_provider_type, **kwargs: synthesizer( + self.type_synthesizer = lambda *args, table_types, **kwargs: synthesizer( *args, **kwargs ) if self.cache: @@ -68,10 +68,10 @@ def __post_init__(self): def __call__( self, *args: TypeOrTypeSynthesizer, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, **kwargs, ) -> TypeOrTypeSynthesizer: - return self.type_synthesizer(*args, offset_provider_type=offset_provider_type, **kwargs) + return self.type_synthesizer(*args, table_types=table_types, **kwargs) TypeOrTypeSynthesizer = Union[ts.TypeSpec, TypeSynthesizer] @@ -313,13 +313,13 @@ def broadcast( def neighbors( offset_literal: it_ts.OffsetLiteralType, it: it_ts.IteratorType, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, ) -> ts.ListType: assert isinstance(offset_literal, it_ts.OffsetLiteralType) and isinstance( offset_literal.value, str ) assert isinstance(it, it_ts.IteratorType) - conn_type = common.get_offset_type(offset_provider_type, offset_literal.value) + conn_type = common.get_offset_type(table_types, offset_literal.value) assert isinstance(conn_type, common.NeighborTableType) return ts.ListType(element_type=it.element_type, offset_type=conn_type.domain[1]) @@ -327,9 +327,7 @@ def neighbors( @_register_builtin_type_synthesizer def lift(stencil: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer - def apply_lift( - *its: it_ts.IteratorType, offset_provider_type: common.TableTypes - ) -> it_ts.IteratorType: + def apply_lift(*its: it_ts.IteratorType, table_types: common.TableTypes) -> it_ts.IteratorType: assert all(isinstance(it, it_ts.IteratorType) for it in its) stencil_args = [ it_ts.IteratorType( @@ -340,7 +338,7 @@ def apply_lift( ) for it in its ] - stencil_return_type = stencil(*stencil_args, offset_provider_type=offset_provider_type) + stencil_return_type = stencil(*stencil_args, table_types=table_types) assert isinstance(stencil_return_type, ts.DataType) position_dims = its[0].position_dims if its else [] @@ -451,7 +449,7 @@ def _canonicalize_nb_fields( def _resolve_dimensions( input_dims: list[common.Dimension], shift_tuple: tuple[itir.OffsetLiteral | itir.CartesianOffset, ...], - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, ) -> list[common.Dimension]: """ Resolves the final dimensions by applying shifts from the given shift tuple. @@ -459,7 +457,7 @@ def _resolve_dimensions( Args: - input_dims: A list of initial dimensions to resolve. - shift_tuple: A tuple of offset literals defining the shift. - - offset_provider_type: Offset provider dictionary. + - table_types: Offset provider dictionary. Returns: A list of resolved dimensions after applying the shifts. @@ -496,11 +494,11 @@ def _resolve_dimensions( ... return common.NeighborTableType( ... connectivity=structure, dtype=None, skip_value=None, max_neighbors=max_neighbors ... ) - >>> offset_provider_type = { + >>> table_types = { ... "C2V": table_type((Cell, C2V), Vertex, 3), ... "V2E": table_type((Vertex, V2E), Edge, 4), ... } - >>> _resolve_dimensions(input_dims, shift_tuple, offset_provider_type) + >>> _resolve_dimensions(input_dims, shift_tuple, table_types) [gt4py.next.iterator.type_system.type_synthesizer.Cell[horizontal], gt4py.next.iterator.type_system.type_synthesizer.K[vertical]] >>> from gt4py.next.iterator.ir_utils import ir_makers as im >>> class IDim(common.CartesianAxisIndex): ... @@ -531,7 +529,7 @@ def _resolve_dimensions( ... ), ... itir.OffsetLiteral(value=0), ... ) - >>> _resolve_dimensions(input_dims, shift_tuple, offset_provider_type) + >>> _resolve_dimensions(input_dims, shift_tuple, table_types) [gt4py.next.iterator.type_system.type_synthesizer.JDim[horizontal], gt4py.next.iterator.type_system.type_synthesizer.IDim[horizontal]] """ @@ -549,7 +547,7 @@ def _resolve_dimensions( assert isinstance(off_literal, itir.OffsetLiteral) and isinstance( off_literal.value, str ) - offset_type = common.get_offset_type(offset_provider_type, off_literal.value) + offset_type = common.get_offset_type(table_types, off_literal.value) if isinstance(offset_type, common.NeighborTableType): if resolved_dim == offset_type.codomain: # Check if input fits to offset resolved_dim = offset_type.domain[0] # Update input_dim for next iteration @@ -566,7 +564,7 @@ def as_fieldop( stencil: TypeSynthesizer, domain: Optional[ts.DomainType] = None, *, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, ) -> TypeSynthesizer: @type_synthesizer def applied_as_fieldop( @@ -588,7 +586,7 @@ def applied_as_fieldop( if not domain: deduced_domain = None output_dims: list[common.Dimension] = [] - if offset_provider_type is not None and shift_sequences_per_param is not None: + if table_types is not None and shift_sequences_per_param is not None: for field, shift_sequences in zip( new_fields, shift_sequences_per_param, strict=True ): @@ -597,9 +595,7 @@ def applied_as_fieldop( for shift_sequence in shift_sequences: output_dims = common.promote_dims( output_dims, - _resolve_dimensions( - input_dims, shift_sequence, offset_provider_type - ), + _resolve_dimensions(input_dims, shift_sequence, table_types), ) assert all(isinstance(dim, common.DimensionMeta) for dim in output_dims) @@ -612,7 +608,7 @@ def applied_as_fieldop( stencil_return = stencil( *(_convert_as_fieldop_input_to_iterator(domain, field) for field in new_fields), - offset_provider_type=offset_provider_type, + table_types=table_types, ) assert isinstance(stencil_return, ts.DataType) @@ -641,10 +637,8 @@ def scan( assert isinstance(direction, ts.ScalarType) and direction.kind == ts.ScalarKind.BOOL @type_synthesizer - def apply_scan( - *its: it_ts.IteratorType, offset_provider_type: common.TableTypes - ) -> ts.DataType: - result = scan_pass(init, *its, offset_provider_type=offset_provider_type) + def apply_scan(*its: it_ts.IteratorType, table_types: common.TableTypes) -> ts.DataType: + result = scan_pass(init, *its, table_types=table_types) assert isinstance(result, ts.DataType) return result @@ -654,11 +648,11 @@ def apply_scan( @_register_builtin_type_synthesizer def map_list(op: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer - def applied_map(*args: ts.ListType, offset_provider_type: common.TableTypes) -> ts.ListType: + def applied_map(*args: ts.ListType, table_types: common.TableTypes) -> ts.ListType: assert len(args) > 0 assert all(isinstance(arg, ts.ListType) for arg in args) arg_el_types = [arg.element_type for arg in args] - el_type = op(*arg_el_types, offset_provider_type=offset_provider_type) + el_type = op(*arg_el_types, table_types=table_types) assert isinstance(el_type, ts.DataType) offset_types = [arg.offset_type for arg in args if arg.offset_type is not None] offset_type = offset_types[0] @@ -675,13 +669,13 @@ def _make_tuple_map_synthesizer( def tuple_map_synthesizer(op: TypeSynthesizer) -> TypeSynthesizer: @type_synthesizer - def applied_map(arg: ts.TupleType, offset_provider_type: common.TableTypes) -> ts.TupleType: + def applied_map(arg: ts.TupleType, table_types: common.TableTypes) -> ts.TupleType: if not isinstance(arg, ts.TupleType): raise TypeError( f"'{builtin_name}' requires a 'TupleType' argument, got '{type(arg).__name__}'." ) - bound_op = functools.partial(op, offset_provider_type=offset_provider_type) + bound_op = functools.partial(op, table_types=table_types) if recursive: return utils.tree_map( # type: ignore[return-value] @@ -709,20 +703,18 @@ def applied_map(arg: ts.TupleType, offset_provider_type: common.TableTypes) -> t @_register_builtin_type_synthesizer def reduce(op: TypeSynthesizer, init: ts.TypeSpec) -> TypeSynthesizer: @type_synthesizer - def applied_reduce(*args: ts.ListType, offset_provider_type: common.TableTypes): + def applied_reduce(*args: ts.ListType, table_types: common.TableTypes): assert all(isinstance(arg, ts.ListType) for arg in args) assert any( arg.offset_type is not None for arg in args ) # we only have `make_const_list`s in the reduce which is not allowed - return op( - init, *(arg.element_type for arg in args), offset_provider_type=offset_provider_type - ) + return op(init, *(arg.element_type for arg in args), table_types=table_types) return applied_reduce @_register_builtin_type_synthesizer -def shift(*offset_literals, offset_provider_type: common.TableTypes) -> TypeSynthesizer: +def shift(*offset_literals, table_types: common.TableTypes) -> TypeSynthesizer: @type_synthesizer def apply_shift( it: it_ts.IteratorType | ts.DeferredType, @@ -733,7 +725,7 @@ def apply_shift( if it.position_dims == "unknown": # nothing to do here return it new_position_dims: list[common.Dimension] | str - if offset_provider_type: + if table_types: new_position_dims = [*it.position_dims] assert len(offset_literals) % 2 == 0 for offset_axis, _ in zip(offset_literals[:-1:2], offset_literals[1::2], strict=True): @@ -744,7 +736,7 @@ def apply_shift( else: assert isinstance(offset_axis, it_ts.OffsetLiteralType) assert isinstance(offset_axis.value, str) - type_ = common.get_offset_type(offset_provider_type, offset_axis.value) + type_ = common.get_offset_type(table_types, offset_axis.value) assert isinstance(type_, common.NeighborTableType) source_dim, target_dim = type_.domain[0], type_.codomain diff --git a/src/gt4py/next/otf/arguments.py b/src/gt4py/next/otf/arguments.py index c0f86a96f6..7c35a00dbd 100644 --- a/src/gt4py/next/otf/arguments.py +++ b/src/gt4py/next/otf/arguments.py @@ -144,7 +144,7 @@ class CompileTimeArgs: argument_descriptor_contexts: ArgStaticDescriptorsContextsByType @property - def offset_provider_type(self) -> common.TableTypes: + def table_types(self) -> common.TableTypes: return common.offset_provider_to_type(self.offset_provider) @classmethod diff --git a/src/gt4py/next/otf/compiled_program.py b/src/gt4py/next/otf/compiled_program.py index f614b0956e..075a3758ba 100644 --- a/src/gt4py/next/otf/compiled_program.py +++ b/src/gt4py/next/otf/compiled_program.py @@ -41,7 +41,8 @@ ScalarOrTupleOfScalars: TypeAlias = xtyping.MaybeNestedInTuple[core_defs.Scalar] -#: Content of the key: (*hashable_arg_descriptors, id(offset_provider), concrete_instantation_if_generic) +#: Content of the key: (*hashable_arg_descriptors, hash of the offset provider's (tag, id(table)) +#: items, concrete_instantation_if_generic) CompiledProgramsKey: TypeAlias = tuple[tuple[Hashable, ...], int, str | None] ArgStaticDescriptorsByType: TypeAlias = dict[ @@ -643,6 +644,7 @@ def _compile_variant( else: raise ValueError(f"Invalid 'offset_provider': {offset_provider}") + common.check_offset_provider(offset_provider, deep=True) self._initialize_argument_descriptor_mapping(argument_descriptors) _validate_argument_descriptors(self.program_type, argument_descriptors) diff --git a/src/gt4py/next/otf/options.py b/src/gt4py/next/otf/options.py index 6f0d7ac7a7..c7863ef319 100644 --- a/src/gt4py/next/otf/options.py +++ b/src/gt4py/next/otf/options.py @@ -15,7 +15,7 @@ class CompilationOptionsArgs(TypedDict, total=False): enable_jit: bool static_params: Sequence[str] - connectivities: common.OffsetProvider + connectivities: common.OffsetProviderLike static_domains: bool @@ -36,9 +36,17 @@ class CompilationOptions: #: A dictionary holding static/compile-time information about the offset providers. #: For now, it is used for ahead of time compilation in DaCe orchestrated programs, #: i.e. DaCe programs that call GT4Py Programs -SDFGConvertible interface-. - connectivities: common.OffsetProvider | None = None + connectivities: common.OffsetProviderLike | None = None static_domains: bool = False + def __post_init__(self) -> None: + if self.connectivities is not None: + object.__setattr__( + self, "connectivities", common.as_tag_keyed_offset_provider(self.connectivities) + ) + # the DaCe orchestration reads these directly, without passing an offset provider + common.check_offset_provider(self.connectivities, deep=True) + assert CompilationOptionsArgs.__annotations__.keys() == CompilationOptions.__annotations__.keys() diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index 3c4f2ce1f2..c80e8aa55b 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -16,7 +16,6 @@ from gt4py._core import definitions as core_defs from gt4py.eve import codegen from gt4py.next import common -from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import ir as itir from gt4py.next.iterator.transforms import pass_manager from gt4py.next.otf import artifacts, stages, workflow @@ -68,7 +67,7 @@ def _process_regular_arguments( self, program: itir.Program, arg_types: tuple[ts.TypeSpec, ...], - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, ) -> tuple[list[interface.Parameter], list[str]]: parameters: list[interface.Parameter] = [] arg_exprs: list[str] = [] @@ -82,21 +81,14 @@ def _process_regular_arguments( if isinstance(parameter.type_, ts.FieldType): for dim in parameter.type_.dims: - if ( - isinstance( - dim, fbuiltins.FieldOffset - ) # TODO(havogt): remove support for FieldOffset as Dimension - or common.is_local_dimension(dim) - ): + if common.is_local_dimension(dim): # translate sparse dimensions to tuple dtype - # NOTE: the tag is the offset-provider key, and its mangled form names the - # `generated::_t` tag type. A legacy `FieldOffset` carries it as `value`. - dim_name = dim.value if isinstance(dim, fbuiltins.FieldOffset) else dim.tag + # NOTE: the local dimension's tag names the `generated::_t` tag type + # (mangled); its table may be keyed by a connectivity sharing it. + dim_name = dim.tag connectivity = common.get_offset_type( - offset_provider_type, - dim_name - if isinstance(dim, fbuiltins.FieldOffset) - else common.connectivity_key_over(offset_provider_type, dim), + table_types, + common.connectivity_key_over(table_types, dim), ) assert isinstance(connectivity, common.NeighborTableType) size = connectivity.max_neighbors @@ -105,12 +97,12 @@ def _process_regular_arguments( return parameters, arg_exprs def _process_connectivity_args( - self, offset_provider_type: common.TableTypes + self, table_types: common.TableTypes ) -> tuple[list[interface.Parameter], list[str]]: parameters: list[interface.Parameter] = [] arg_exprs: list[str] = [] - for name, connectivity_type in offset_provider_type.items(): + for name, connectivity_type in table_types.items(): if isinstance(connectivity_type, common.NeighborTableType): if connectivity_type.dtype.scalar_type not in [np.int32, np.int64]: raise ValueError( @@ -181,7 +173,7 @@ def generate_stencil_source( gtfn_ir = GTFN_lowering.apply( new_program, - offset_provider_type=common.offset_provider_to_type(offset_provider), + table_types=common.offset_provider_to_type(offset_provider), column_axis=column_axis, ) @@ -198,13 +190,13 @@ def __call__( # the program) arg_types = inp.args.args regular_parameters, regular_args_expr = self._process_regular_arguments( - program, arg_types, inp.args.offset_provider_type + program, arg_types, inp.args.table_types ) # handle connectivity parameters and arguments (i.e. what the user provided in the offset # provider) connectivity_parameters, connectivity_args_expr = self._process_connectivity_args( - inp.args.offset_provider_type + inp.args.table_types ) # combine into a format that is aligned with what the backend expects diff --git a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py index c893be66e9..69d44bbfac 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/itir_to_gtfn_ir.py @@ -159,10 +159,10 @@ def _collect_dimensions_from_params( def _collect_offset_definitions( node: itir.Node, grid_type: common.GridType, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, ) -> dict[str, TagDefinition]: offset_definitions = {} - offset_provider_type = {**offset_provider_type} + table_types = {**table_types} cartesian_offsets: OrderedSet[itir.CartesianOffset] = OrderedSet( node.walk_values().if_isinstance(itir.CartesianOffset).to_list() @@ -187,7 +187,7 @@ def _collect_offset_definitions( name=Sym(id=common.codegen_name(dim.tag)), alias=_vertical_dimension ) - for offset_name, connectivity_type in offset_provider_type.items(): + for offset_name, connectivity_type in table_types.items(): if isinstance(connectivity_type, common.NeighborTableType): assert grid_type == common.GridType.UNSTRUCTURED offset_definitions[offset_name] = TagDefinition( @@ -339,7 +339,7 @@ class GTFN_lowering(eve.NodeTranslator, eve.VisitorWithSymbolTableTrait): } _unary_op_map: ClassVar[dict[str, str]] = {"not_": "!"} - offset_provider_type: common.TableTypes + table_types: common.TableTypes column_axis: Optional[common.Dimension] grid_type: common.GridType @@ -354,19 +354,19 @@ def apply( cls, node: itir.Program, *, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, column_axis: Optional[common.Dimension], ) -> Program: if not isinstance(node, itir.Program): raise TypeError(f"Expected a 'Program', got '{type(node).__name__}'.") - node = itir_type_inference.infer(node, offset_provider_type=offset_provider_type) + node = itir_type_inference.infer(node, table_types=table_types) grid_type = ir_utils_misc.grid_type_from_program(node) if grid_type == common.GridType.UNSTRUCTURED: node = _CannonicalizeUnstructuredDomain.apply(node) - return cls( - offset_provider_type=offset_provider_type, column_axis=column_axis, grid_type=grid_type - ).visit(node) + return cls(table_types=table_types, column_axis=column_axis, grid_type=grid_type).visit( + node + ) def visit_Sym(self, node: itir.Sym, **kwargs: Any) -> Sym: return Sym(id=node.id) @@ -504,8 +504,8 @@ def _visit_unstructured_domain(self, node: itir.FunCall, **kwargs: Any) -> Node: if "stencil" in kwargs: shift_offsets = self._collect_offset_or_axis_node(itir.OffsetLiteral, kwargs["stencil"]) for o in shift_offsets: - if o in self.offset_provider_type and isinstance( - common.get_offset_type(self.offset_provider_type, o), + if o in self.table_types and isinstance( + common.get_offset_type(self.table_types, o), common.NeighborTableType, ): # `o` is an offset-provider key, i.e. a qualified tag: mangle it exactly as @@ -721,7 +721,7 @@ def visit_Program(self, node: itir.Program, **kwargs: Any) -> Program: offset_definitions = { **_collect_dimensions_from_params(node.params, self.grid_type), **_collect_dimensions_from_domain(node.body), - **_collect_offset_definitions(node, self.grid_type, self.offset_provider_type), + **_collect_offset_definitions(node, self.grid_type, self.table_types), } offset_definitions = _add_staggered_aliases(offset_definitions) return Program( diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index f195670437..69beedd2a4 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -87,7 +87,7 @@ class DataflowBuilder(Protocol): """Visitor interface to build a dataflow subgraph.""" @abc.abstractmethod - def get_offset_provider_type(self, offset: str) -> gtx_common.NeighborTableType: ... + def get_table_type(self, offset: str) -> gtx_common.NeighborTableType: ... @abc.abstractmethod def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: @@ -560,17 +560,17 @@ class GTIRToSDFG(eve.NodeVisitor, SDFGBuilder): from where to continue building the SDFG. """ - offset_provider_type: gtx_common.TableTypes + table_types: gtx_common.TableTypes column_axis: Optional[gtx_common.Dimension] uids: gtx_utils.IDGeneratorPool = dataclasses.field( init=False, repr=False, default_factory=lambda: gtx_utils.IDGeneratorPool() ) - def get_offset_provider_type(self, offset: str) -> gtx_common.NeighborTableType: - return gtx_common.get_offset_type(self.offset_provider_type, offset) + def get_table_type(self, offset: str) -> gtx_common.NeighborTableType: + return gtx_common.get_offset_type(self.table_types, offset) def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: - return gtx_common.connectivity_key_over(self.offset_provider_type, local_dim) + return gtx_common.connectivity_key_over(self.table_types, local_dim) def make_field( self, @@ -730,7 +730,7 @@ def add_nested_sdfg( connectivity_arrays = { gtx_dace_args.connectivity_identifier(offset) - for offset in gtx_dace_args.filter_connectivity_types(self.offset_provider_type) + for offset in gtx_dace_args.filter_connectivity_types(self.table_types) } inner_ctx_globals = [ @@ -844,13 +844,13 @@ def _make_array_shape_and_strides( Returns: Two lists of symbols, one for the shape and the other for the strides of the array. """ - neighbor_table_types = gtx_dace_args.filter_connectivity_types(self.offset_provider_type) + neighbor_table_types = gtx_dace_args.filter_connectivity_types(self.table_types) shape = [] for dim in dims: if gtx_common.is_local_dimension(dim): # for local dimension, the size is taken from the associated connectivity type shape.append(gtx_dace_args.local_dimension_size(name, dim, neighbor_table_types)) - elif gtx_dace_args.is_connectivity_identifier(name, self.offset_provider_type): + elif gtx_dace_args.is_connectivity_identifier(name, self.table_types): # we use symbolic size for the global dimension of a connectivity shape.append(gtx_dace_args.field_size_symbol(name, dim, neighbor_table_types)) else: @@ -1037,7 +1037,7 @@ def _add_sdfg_params( # add SDFG storage for connectivity tables for offset, connectivity_type in gtx_dace_args.filter_connectivity_types( - self.offset_provider_type + self.table_types ).items(): gt_type = ts.FieldType( dims=[connectivity_type.domain[0], connectivity_type.domain[1]], @@ -1095,7 +1095,7 @@ def visit_Program(self, node: gtir.Program) -> dace.SDFG: unused_connectivities = [ data for data, datadesc in nsdfg.arrays.items() - if gtx_dace_args.is_connectivity_identifier(data, self.offset_provider_type) + if gtx_dace_args.is_connectivity_identifier(data, self.table_types) and datadesc.transient ] for data in unused_connectivities: @@ -1396,7 +1396,7 @@ def visit_SymRef( def lower_program_to_sdfg( ir: gtir.Program, - offset_provider_type: gtx_common.TableTypes, + table_types: gtx_common.TableTypes, column_axis: Optional[gtx_common.Dimension] = None, ) -> dace.SDFG: """ @@ -1407,7 +1407,7 @@ def lower_program_to_sdfg( Args: ir: The GTIR program node to be lowered to SDFG - offset_provider_type: The definitions of offset providers used by the program node + table_types: The definitions of offset providers used by the program node column_axis: Vertical dimension used for column scan expressions. Returns: @@ -1419,7 +1419,7 @@ def lower_program_to_sdfg( if ir.declarations: raise NotImplementedError("Temporaries not supported yet by GTIR DaCe backend.") - ir = gtir_type_inference.infer(ir, offset_provider_type=offset_provider_type) + ir = gtir_type_inference.infer(ir, table_types=table_types) ir = ir_prune_casts.PruneCasts().visit(ir) # DaCe requires C-compatible strings for the names of data containers, @@ -1428,7 +1428,7 @@ def lower_program_to_sdfg( # Here we find new names for invalid symbols present in the IR. ir = gtir_to_sdfg_utils.replace_invalid_symbols(ir) - sdfg_genenerator = GTIRToSDFG(offset_provider_type, column_axis) + sdfg_genenerator = GTIRToSDFG(table_types, column_axis) sdfg = sdfg_genenerator.visit(ir) assert isinstance(sdfg, dace.SDFG) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py index 0eb550936c..0ffa95e4b1 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py @@ -254,12 +254,10 @@ def translate_concat_where( local_dim = node.type.dtype.offset_type assert local_dim is not None dtype = gtx_dace_args.as_dace_type(node.type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type( - sdfg_builder.connectivity_key_over(local_dim) - ) - assert isinstance(offset_provider_type, gtx_common.NeighborTableType) + table_type = sdfg_builder.get_table_type(sdfg_builder.connectivity_key_over(local_dim)) + assert isinstance(table_type, gtx_common.NeighborTableType) local_idx = gtx_common.order_dimensions([*output_dims, local_dim]).index(local_dim) - output_shape.insert(local_idx, offset_provider_type.max_neighbors) + output_shape.insert(local_idx, table_type.max_neighbors) output, output_desc = sdfg_builder.add_temp_array(ctx.sdfg, output_shape, dtype) output_node = ctx.state.add_access(output) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 76cbfeb05b..94d52f94f5 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py @@ -762,7 +762,7 @@ def _visit_if_branch_arg( local_dim = arg.gt_dtype.offset_type assert local_dim is not None assert isinstance( - self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.get_table_type( self.subgraph_builder.connectivity_key_over(local_dim) ), gtx_common.NeighborTableType, @@ -1078,7 +1078,7 @@ def _visit_neighbors(self, node: gtir.FunCall) -> ValueExpr: assert isinstance(node.args[0], gtir.OffsetLiteral) offset = node.args[0].value assert isinstance(offset, str) - conn_type = self.subgraph_builder.get_offset_provider_type(offset) + conn_type = self.subgraph_builder.get_table_type(offset) assert isinstance(conn_type, gtx_common.NeighborTableType) it = self.visit(node.args[1]) @@ -1317,7 +1317,7 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if offset_type is gtx_common.ConstList: # this input argument is the result of `make_const_list` continue - offset_provider_t = self.subgraph_builder.get_offset_provider_type( + offset_provider_t = self.subgraph_builder.get_table_type( self.subgraph_builder.connectivity_key_over(offset_type) ) assert isinstance(offset_provider_t, gtx_common.NeighborTableType) @@ -1437,7 +1437,7 @@ def _broadcast_const_list( self, const_list: MemletExpr | ValueExpr, list_type: ts.ListType ) -> ValueExpr: assert list_type.offset_type is not None - offset_provider_t = self.subgraph_builder.get_offset_provider_type( + offset_provider_t = self.subgraph_builder.get_table_type( self.subgraph_builder.connectivity_key_over(list_type.offset_type) ) assert isinstance(offset_provider_t, gtx_common.NeighborTableType) @@ -1478,15 +1478,15 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: and input_expr.gt_dtype.offset_type is not None ) offset_type = input_expr.gt_dtype.offset_type - offset_provider_type = self.subgraph_builder.get_offset_provider_type( + table_type = self.subgraph_builder.get_table_type( self.subgraph_builder.connectivity_key_over(offset_type) ) - assert isinstance(offset_provider_type, gtx_common.NeighborTableType) + assert isinstance(table_type, gtx_common.NeighborTableType) inp_conn = "_in" outp_conn = "_out" mask_conn = "_mask" - if offset_provider_type.has_skip_values: + if table_type.has_skip_values: assert ( isinstance(input_expr.gt_dtype, ts.ListType) and input_expr.gt_dtype.offset_type is not None @@ -1509,12 +1509,10 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: ) self.state.add_node(reduce_node) - origin_map_index = gtir_to_sdfg_utils.get_map_variable(offset_provider_type.domain[0]) + origin_map_index = gtir_to_sdfg_utils.get_map_variable(table_type.domain[0]) self._add_input_data_edge( self.state.add_access(connectivity), - dace_subsets.Range.from_string( - f"{origin_map_index}, 0:{offset_provider_type.max_neighbors}" - ), + dace_subsets.Range.from_string(f"{origin_map_index}, 0:{table_type.max_neighbors}"), reduce_node, mask_conn, ) @@ -1758,10 +1756,8 @@ def _visit_shift(self, node: gtir.FunCall) -> IteratorExpr: else: assert isinstance(offset_provider_arg, gtir.OffsetLiteral) assert isinstance(offset_provider_arg.value, str) - offset_provider_type = self.subgraph_builder.get_offset_provider_type( - offset_provider_arg.value - ) - assert isinstance(offset_provider_type, gtx_common.NeighborTableType) + table_type = self.subgraph_builder.get_table_type(offset_provider_arg.value) + assert isinstance(table_type, gtx_common.NeighborTableType) # a named offset → unstructured shift; the offset value may be a static # `OffsetLiteral` or a dynamic offset (handled by `_make_unstructured_shift`). # initially, the storage for the connectivity tables is created as transient; @@ -1771,9 +1767,7 @@ def _visit_shift(self, node: gtir.FunCall) -> IteratorExpr: self.sdfg.arrays[offset_table].transient = False offset_table_node = self.state.add_access(offset_table) - return self._make_unstructured_shift( - it, offset_provider_type, offset_table_node, offset_expr - ) + return self._make_unstructured_shift(it, table_type, offset_table_node, offset_expr) def _visit_generic_builtin(self, node: gtir.FunCall) -> ValueExpr: """ diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index d779e71fef..cd3805ec52 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -325,11 +325,11 @@ def _construct_if_branch_output( assert out_type.dtype.offset_type is not None assert isinstance(out_type.dtype.element_type, ts.ScalarType) dtype = gtx_dace_args.as_dace_type(out_type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type( + table_type = sdfg_builder.get_table_type( sdfg_builder.connectivity_key_over(out_type.dtype.offset_type) ) - assert isinstance(offset_provider_type, gtx_common.NeighborTableType) - shape = [*shape, offset_provider_type.max_neighbors] + assert isinstance(table_type, gtx_common.NeighborTableType) + shape = [*shape, table_type.max_neighbors] out, _ = sdfg_builder.add_temp_array(ctx.sdfg, shape, dtype) out_node = ctx.state.add_access(out) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index 911078c9c7..4a4f06b692 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -384,11 +384,9 @@ def get_scan_output_shape( assert isinstance(scan_init_data.gt_type, ts.ListType) assert scan_init_data.gt_type.offset_type offset_type = scan_init_data.gt_type.offset_type - offset_provider_type = sdfg_builder.get_offset_provider_type( - sdfg_builder.connectivity_key_over(offset_type) - ) - assert isinstance(offset_provider_type, gtx_common.NeighborTableType) - list_size = offset_provider_type.max_neighbors + table_type = sdfg_builder.get_table_type(sdfg_builder.connectivity_key_over(offset_type)) + assert isinstance(table_type, gtx_common.NeighborTableType) + list_size = table_type.max_neighbors return [scan_column_size, dace.symbolic.SymExpr(list_size)] if isinstance(init_data, tuple): diff --git a/src/gt4py/next/program_processors/runners/dace/program.py b/src/gt4py/next/program_processors/runners/dace/program.py index f8f8c7e1a4..29b0764619 100644 --- a/src/gt4py/next/program_processors/runners/dace/program.py +++ b/src/gt4py/next/program_processors/runners/dace/program.py @@ -65,7 +65,7 @@ def __sdfg__(self, *args: Any, **kwargs: Any) -> dace.sdfg.sdfg.SDFG: program, ) object.__setattr__( - gtir_stage.args, "offset_provider", gtir_stage.args.offset_provider_type + gtir_stage.args, "offset_provider", gtir_stage.args.table_types ) # TODO(ricoh): currently this is circumventing the frozenness of CompileTimeArgs # in order to isolate DaCe from the runtime tables in connectivities.offset_provider. # These are needed at the time of writing for mandatory GTIR passes. @@ -167,27 +167,21 @@ def __sdfg_closure__(self, reevaluate: Optional[dict[str, str]] = None) -> dict[ # Build the closure dictionary closure_dict: dict[str, dace.data.Array] = {} - offset_provider_type = gtx_common.offset_provider_to_type( - self.compilation_options.connectivities - ) + table_types = gtx_common.offset_provider_to_type(self.compilation_options.connectivities) for conn_id, conn in used_connectivities.items(): if conn_id not in self.connectivity_tables_data_descriptors: self.connectivity_tables_data_descriptors[conn_id] = dace.data.Array( dtype=dace.dtypes.dtype_to_typeclass(conn.dtype.dtype.type), shape=[ - gtx_dace_args.field_size_symbol( - conn_id, conn.domain.dims[0], offset_provider_type - ), - gtx_dace_args.field_size_symbol( - conn_id, conn.domain.dims[1], offset_provider_type - ), + gtx_dace_args.field_size_symbol(conn_id, conn.domain.dims[0], table_types), + gtx_dace_args.field_size_symbol(conn_id, conn.domain.dims[1], table_types), ], strides=[ gtx_dace_args.field_stride_symbol( - conn_id, conn.domain.dims[0], offset_provider_type + conn_id, conn.domain.dims[0], table_types ), gtx_dace_args.field_stride_symbol( - conn_id, conn.domain.dims[1], offset_provider_type + conn_id, conn.domain.dims[1], table_types ), ], storage=Program.connectivity_tables_data_descriptors["storage"], diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py index 55d3c77ada..d0308b5f03 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -59,32 +59,30 @@ def connectivity_identifier(name: str) -> str: return f"{CONNECTIVITY_INDENTIFIER_PREFIX}{gtx_common.codegen_name(name)}" -def is_connectivity_identifier( - name: str, offset_provider_type: gtx_common.TableTypes | None = None -) -> bool: +def is_connectivity_identifier(name: str, table_types: gtx_common.TableTypes | None = None) -> bool: if (m := CONNECTIVITY_INDENTIFIER_RE.match(name)) is None: return False - elif offset_provider_type is None: + elif table_types is None: # If no offset provider type is provided, we assume there is a connectivity identifier # that matches the CONNECTIVITY_INDENTIFIER_RE. return True else: - return gtx_common.has_offset(offset_provider_type, gtx_common.from_codegen_name(m[1])) + return gtx_common.has_offset(table_types, gtx_common.from_codegen_name(m[1])) def _field_symbol( field_name: str, dim: gtx_common.Dimension, sym: Literal["size", "stride"], - offset_provider_type: gtx_common.TableTypes | None, + table_types: gtx_common.TableTypes | None, ) -> dace.symbol: if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is None: name = f"__{field_name}_{gtx_common.codegen_name(dim.tag)}_{sym}" else: # a connectivity field - assert offset_provider_type is not None + assert table_types is not None offset = gtx_common.from_codegen_name(m[1]) - assert offset in offset_provider_type - conn_type = offset_provider_type[offset] + assert offset in table_types + conn_type = table_types[offset] assert isinstance(conn_type, gtx_common.NeighborTableType) if dim == conn_type.domain[0]: name = f"__{field_name}_source_{sym}" @@ -98,17 +96,17 @@ def _field_symbol( def field_size_symbol( field_name: str, dim: gtx_common.Dimension, - offset_provider_type: gtx_common.TableTypes, + table_types: gtx_common.TableTypes, ) -> dace.symbol: - return _field_symbol(field_name, dim, "size", offset_provider_type) + return _field_symbol(field_name, dim, "size", table_types) def field_stride_symbol( field_name: str, dim: gtx_common.Dimension, - offset_provider_type: gtx_common.TableTypes | None = None, + table_types: gtx_common.TableTypes | None = None, ) -> dace.symbol: - return _field_symbol(field_name, dim, "stride", offset_provider_type) + return _field_symbol(field_name, dim, "stride", table_types) def local_dimension_size( @@ -151,7 +149,7 @@ def range_stop_symbol(field_name: str, dim: gtx_common.Dimension) -> dace.symbol def filter_connectivity_types( - offset_provider_type: gtx_common.TableTypes, + table_types: gtx_common.TableTypes, ) -> dict[str, gtx_common.NeighborTableType]: """ Filter offset provider types of type `NeighborTableType`. @@ -160,6 +158,6 @@ def filter_connectivity_types( """ return { offset: conn - for offset, conn in offset_provider_type.items() + for offset, conn in table_types.items() if isinstance(conn, gtx_common.NeighborTableType) } diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py index 1af014470c..d21c475723 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_callable.py @@ -92,14 +92,17 @@ def _get_args(sdfg: dace.SDFG, args: Sequence[Any]) -> dict[str, Any]: def get_sdfg_conn_args( sdfg: dace.SDFG, - offset_provider: gtx_common.OffsetProvider, + offset_provider: gtx_common.OffsetProviderLike, ) -> dict[str, core_defs.NDArrayObject]: """ Extracts the connectivity tables that are used in the sdfg and ensures that the memory buffers are allocated for the target device. """ connectivity_args = {} - for offset, connectivity in offset_provider.items(): + # NOTE: not strict, like the other IR-level hooks: the keys of a hand-written program are its + # own business, and a declaration is normalized to its tag either way. + provider = gtx_common.as_tag_keyed_offset_provider(offset_provider, strict=False) + for offset, connectivity in provider.items(): name = gtx_dace_args.connectivity_identifier(offset) if name in sdfg.arrays: assert gtx_common.is_neighbor_table(connectivity) diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index df2ceb6670..bdb379e9f2 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -31,7 +31,7 @@ def find_constant_symbols( ir: itir.Program, sdfg: dace.SDFG, - offset_provider_type: common.TableTypes, + table_types: common.TableTypes, disable_field_origin_on_program_arguments: bool, unstructured_horizontal_has_unit_stride: bool, ) -> dict[str, int]: @@ -54,7 +54,7 @@ def find_constant_symbols( sdfg_stride_symbol = gtx_dace_args.field_stride_symbol(str(p.id), dim) constant_symbols[sdfg_stride_symbol.name] = 1 # Same for connectivity tables, for which the first dimension is always horizontal - for offset, conn_type in offset_provider_type.items(): + for offset, conn_type in table_types.items(): if ( isinstance(conn_type, common.NeighborTableType) and (conn_id := gtx_dace_args.connectivity_identifier(offset)) in sdfg.arrays @@ -62,7 +62,7 @@ def find_constant_symbols( assert not sdfg.arrays[conn_id].transient assert conn_type.domain[0].kind == common.DimensionKind.HORIZONTAL sdfg_stride_symbol = gtx_dace_args.field_stride_symbol( - conn_id, conn_type.domain[0], offset_provider_type + conn_id, conn_type.domain[0], table_types ) constant_symbols[sdfg_stride_symbol.name] = 1 @@ -396,15 +396,15 @@ def _generate_sdfg_without_configuring_dace( use_max_domain_range_on_unstructured_shift=self.use_max_domain_range_on_unstructured_shift, offset_provider=offset_provider, ) - offset_provider_type = common.offset_provider_to_type(offset_provider) + table_types = common.offset_provider_to_type(offset_provider) on_gpu = self.device_type != core_defs.DeviceType.CPU - sdfg = gtx_dace_lowering.lower_program_to_sdfg(ir, offset_provider_type, column_axis) + sdfg = gtx_dace_lowering.lower_program_to_sdfg(ir, table_types, column_axis) constant_symbols = find_constant_symbols( ir, sdfg, - offset_provider_type, + table_types, self.disable_field_origin_on_program_arguments, self.unstructured_horizontal_has_unit_stride, ) @@ -468,7 +468,7 @@ def __call__( sdfg = self.generate_sdfg( program, - inp.args.offset_provider, # TODO(havogt): should be offset_provider_type once the transformation don't require run-time info + inp.args.offset_provider, # TODO(havogt): should be table_types once the transformation don't require run-time info inp.args.column_axis, ) diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 648dfa866b..5f0a20ca5c 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -71,9 +71,10 @@ class ShiftType(TypeSpec): """ The type of a shift: it takes a field over `codomain` to a field over `domain`. - `domain` has one dimension for a Cartesian shift (`KDim + 1`) and for a single neighbor - (`V2E[i]`), and two -- the connectivity's domain and its local dimension -- for all - neighbors (`V2E`). + `domain` has one dimension for a Cartesian shift (`KDim + 1`, `as_offset(KDim, offsets)`) + and for a single neighbor (`V2E[i]`), and two -- the connectivity's domain and its local + dimension -- for all neighbors (`V2E`). The type of the *table* bound to a connectivity is a + `common.NeighborTableType`. """ codomain: common.Dimension diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 9fc0f7c1ae..0306f8e2f5 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 0029) that would be two different dimensions, and tests that mix a # `toy_connectivity` connectivity with a `cases_utils` mesh would silently stop matching. -from next_tests.toy_connectivity import C2EDim, Cell, E2VDim, Edge, V2EDim, Vertex +from next_tests.toy_connectivity import ( + C2E, + C2EDim, + Cell, + E2V, + E2VDim, + Edge, + V2E, + V2EDim, + Vertex, +) __all__ = [ @@ -52,7 +62,6 @@ "Vertex", "Edge", "Cell", - "EdgeOffset", "MeshDescriptor", "CartesianGridDescriptor", ] @@ -177,20 +186,12 @@ class KDim(gtx.CartesianAxisIndex, kind=gtx.DimensionKind.VERTICAL): ... KHalfDim = common.flip_staggered(KDim) -Ioff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) -Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) - - -EdgeOffset = gtx.FieldOffset("EdgeOffset", source=Edge, target=(Edge,)) - -class C2VDim(gtx.LocalDimensionIndex): ... +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 @@ -206,7 +207,7 @@ def sizes(self) -> tuple[int, int, int]: ... def offset_provider(self) -> common.OffsetProvider: ... @property - def offset_provider_type(self) -> common.TableTypes: ... + def table_types(self) -> common.TableTypes: ... def simple_cartesian_grid( @@ -222,7 +223,10 @@ def simple_cartesian_grid( name="simple_cartesian_grid", sizes=sizes, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + table_types=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -248,7 +252,7 @@ def num_edges(self) -> int: ... def offset_provider(self) -> common.OffsetProvider: ... @property - def offset_provider_type(self) -> common.TableTypes: ... + def table_types(self) -> common.TableTypes: ... def simple_mesh(allocator) -> MeshDescriptor: @@ -308,28 +312,28 @@ def simple_mesh(allocator) -> MeshDescriptor: e2v_arr = np.asarray(e2v_arr, dtype=gtx.IndexType) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 4}, codomain=Edge, data=v2e_arr, skip_value=None, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 4}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 4}, codomain=Edge, data=c2e_arr, @@ -344,7 +348,10 @@ def simple_mesh(allocator) -> MeshDescriptor: num_edges=np.int32(num_edges), num_cells=num_cells, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + table_types=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -403,28 +410,28 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) offset_provider = { - V2E.value: constructors.as_connectivity( + V2E: constructors.as_connectivity( domain={Vertex: v2e_arr.shape[0], V2EDim: 5}, codomain=Edge, data=v2e_arr, skip_value=common._DEFAULT_SKIP_VALUE, allocator=allocator, ), - E2V.value: constructors.as_connectivity( + E2V: constructors.as_connectivity( domain={Edge: e2v_arr.shape[0], E2VDim: 2}, codomain=Vertex, data=e2v_arr, skip_value=None, allocator=allocator, ), - C2V.value: constructors.as_connectivity( + C2V: constructors.as_connectivity( domain={Cell: c2v_arr.shape[0], C2VDim: 3}, codomain=Vertex, data=c2v_arr, skip_value=None, allocator=allocator, ), - C2E.value: constructors.as_connectivity( + C2E: constructors.as_connectivity( domain={Cell: c2e_arr.shape[0], C2EDim: 3}, codomain=Edge, data=c2e_arr, @@ -439,7 +446,10 @@ def skip_value_mesh(allocator) -> MeshDescriptor: num_edges=num_edges, num_cells=num_cells, offset_provider=offset_provider, - offset_provider_type=common.offset_provider_to_type(offset_provider), + # NOTE: tag-keyed, the form the IR-level APIs (type inference, transformations) take. + table_types=common.as_tag_keyed_offset_provider( + common.offset_provider_to_type(offset_provider) + ), ) @@ -452,3 +462,18 @@ def skip_value_mesh(allocator) -> MeshDescriptor: ) def mesh_descriptor(request, exec_alloc_descriptor) -> MeshDescriptor: yield request.param(exec_alloc_descriptor.allocator) + + +def ir_level(mesh: MeshDescriptor) -> MeshDescriptor: + """ + A copy of `mesh` whose offset provider is keyed by tags, as the IR-level APIs expect. + + User-facing entry points normalize a class-keyed provider themselves; tests that drive the + lowering or the backends directly have to hand them the tag-keyed form. + """ + return types.SimpleNamespace( + **{ + **vars(mesh), + "offset_provider": common.as_tag_keyed_offset_provider(mesh.offset_provider), + } + ) diff --git a/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py b/tests/next_tests/integration_tests/feature_tests/dace_tests/test_write_back_buffer_elimination_lowering.py index a510f383ce..b143f9b113 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.LocalDimensionIndex): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -C2E = gtx.FieldOffset(C2EDim.tag, source=Edge, target=(Cell, C2EDim)) +C2EDim = C2E.Local C2E_TABLE = np.array( [ diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py index 1e65cd734c..c156a61407 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_cartesian_shifts.py @@ -19,8 +19,6 @@ cartesian_case, ) from next_tests.integration_tests.cases_utils import ( - Ioff, - Koff, exec_alloc_descriptor, ) @@ -62,10 +60,10 @@ def test_offset_field(cartesian_case): @gtx.field_operator def testee(a: cases.IKField, offset_field: cases.IKField) -> gtx.Field[[IDim, KDim], bool]: - a_i = a(as_offset(Ioff, offset_field)) + a_i = a(as_offset(IDim, offset_field)) # note: this leads to an access to offset_field in # IDim: (0, out.size[I]), KDim: (0, out.size[K]+1) - a_i_k = a_i(as_offset(Koff, offset_field)) + a_i_k = a_i(as_offset(KDim, offset_field)) b_i = a(IDim + 1) b_i_k = b_i(KDim + 1) return a_i_k == b_i_k @@ -97,7 +95,7 @@ def test_offset_field_of_chained_ops(cartesian_case): def testee(a: cases.IKField, offset_field: cases.IKField) -> cases.IKField: b = a + 1 c = b * 2 - return c(as_offset(Koff, offset_field)) + return c(as_offset(KDim, offset_field)) out = cases.allocate(cartesian_case, testee, cases.RETURN)() a = cases.allocate(cartesian_case, testee, "a").extend({KDim: (0, 1)})() diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py index 4a642f5cc4..68c2dcf72b 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_compiled_program.py @@ -252,7 +252,7 @@ def test_compile_unstructured(unstructured_case, compile_testee_unstructured): compile_testee_unstructured(*args, offset_provider=unstructured_case.offset_provider, **kwargs) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), @@ -317,7 +317,7 @@ def test_compile_unstructured_for_two_offset_providers( *args, offset_provider=unstructured_case.offset_provider, **kwargs ) - v2e_numpy = unstructured_case.offset_provider[V2E.value].asnumpy() + v2e_numpy = unstructured_case.offset_provider[V2E].asnumpy() assert np.allclose( kwargs["out"].asnumpy(), np.sum(np.where(v2e_numpy != -1, args[0].asnumpy()[v2e_numpy], 0), axis=1), diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py index dfe013e45a..365de7d473 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_concat_where.py @@ -454,7 +454,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: t = concat_where(Vertex < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -477,7 +477,7 @@ def testee( t = concat_where(KDim < 2, a(V2E), b(V2E)) return neighbor_sum(t, axis=V2EDim) - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() k_mask = np.arange(unstructured_case_3d.default_sizes[KDim]) < 2 cases.verify_with_default_data( unstructured_case_3d, @@ -497,7 +497,7 @@ def test_with_local_and_nonlocal_field(unstructured_case, static_domains: bool): def testee(a: cases.EField, b: cases.VField) -> cases.VField: return neighbor_sum(concat_where(Vertex < 2, a(V2E), b), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -524,7 +524,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), b(V2E)), (c(V2E), d(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -557,7 +557,7 @@ def testee(a: cases.EField) -> tuple[cases.VField, cases.VField]: neighbor_sum(concat_where(Vertex < 2, 3, a(V2E)), axis=V2EDim), ) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, @@ -590,7 +590,7 @@ def testee( t = concat_where(Vertex < 2, (a(V2E), c), (3, b(V2E))) return neighbor_sum(t[0], axis=V2EDim), neighbor_sum(t[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py index 753015f0d7..b65a16ffea 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_external_local_field.py @@ -31,11 +31,11 @@ def testee( ) # multiplication with shifted `ones` because reduction of only non-shifted field with local dimension is not supported inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) ones = cases.allocate(unstructured_case, testee, "ones").strategy(cases.ConstInitializer(1))() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify( unstructured_case, testee, @@ -55,7 +55,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return inp[V2EDim(0)] + inp[V2EDim(1)] + inp[V2EDim(2)] + inp[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -77,7 +77,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int64 return inp_64[V2EDim(0)] + inp_64[V2EDim(1)] + inp_64[V2EDim(2)] + inp_64[V2EDim(3)] inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -99,7 +99,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 return neighbor_sum(inp, axis=V2EDim) inp = unstructured_case.as_field( - [Vertex, V2EDim], unstructured_case.offset_provider[V2EDim.tag].asnumpy() + [Vertex, V2EDim], unstructured_case.offset_provider[V2E].asnumpy() ) cases.verify( @@ -107,7 +107,7 @@ def testee(inp: gtx.Field[[Vertex, V2EDim], int32]) -> gtx.Field[[Vertex], int32 testee, inp, out=cases.allocate(unstructured_case, testee, cases.RETURN)(), - ref=np.sum(unstructured_case.offset_provider[V2EDim.tag].asnumpy(), axis=1), + ref=np.sum(unstructured_case.offset_provider[V2E].asnumpy(), axis=1), ) @@ -119,7 +119,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: return inp(V2E) out = unstructured_case.as_field( - [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider[V2EDim.tag].asnumpy()) + [Vertex, V2EDim], np.zeros_like(unstructured_case.offset_provider[V2E].asnumpy()) ) inp = cases.allocate(unstructured_case, testee, "inp")() cases.verify( @@ -127,5 +127,5 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex, V2EDim], int32]: testee, inp, out=out, - ref=inp.asnumpy()[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], + ref=inp.asnumpy()[unstructured_case.offset_provider[V2E].asnumpy()], ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py index 262751f234..dbc623f2de 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py @@ -60,7 +60,7 @@ def test_import_offset_module_unstructured_shift(unstructured_case): def testee(a: cases.EField) -> cases.VField: return neighbor_sum(a(cases.V2E), axis=cases.V2EDim) - v2e_table = unstructured_case.offset_provider[cases.V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[cases.V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -77,7 +77,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider[cases.E2VDim.tag].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[cases.E2V].asnumpy()[:, 0]], ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py index 39aee5de06..6ceead30d6 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_named_collections.py @@ -483,7 +483,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) ) @@ -641,7 +641,7 @@ def testee( ) return neighbor_sum(t.neighbors, axis=V2EDim), t.center - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() vertex_mask = np.arange(unstructured_case.default_sizes[Vertex]) < 2 cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py index bce54f4e6a..07793c44fe 100644 --- 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 @@ -42,7 +42,7 @@ class V2EShared(gtx.NeighborConnectivity[V, E]): @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() + v2e_arr = mesh.offset_provider[cases_utils.V2E].asnumpy() table = constructors.as_connectivity( domain={V: v2e_arr.shape[0], V2E.Local: v2e_arr.shape[1]}, codomain=E, @@ -65,9 +65,7 @@ def case(exec_alloc_descriptor): 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}, + offset_provider={V2E: table, V2EShared: shared_table}, default_sizes={V: mesh.num_vertices, E: mesh.num_edges, V2E.Local: v2e_arr.shape[1]}, grid_type=common.GridType.UNSTRUCTURED, allocator=exec_alloc_descriptor.allocator, @@ -75,7 +73,7 @@ def case(exec_alloc_descriptor): def _table(case: cases.Case, connectivity=V2E) -> np.ndarray: - return case.offset_provider[connectivity.offset_tag].asnumpy() + return case.offset_provider[connectivity].asnumpy() @pytest.mark.uses_unstructured_shift @@ -148,9 +146,7 @@ def testee( @pytest.fixture def case_without_owner(case): """Only the sharing connectivity is bound: enough for a shift, which needs only its table.""" - return dataclasses.replace( - case, offset_provider={V2EShared.offset_tag: case.offset_provider[V2EShared.offset_tag]} - ) + return dataclasses.replace(case, offset_provider={V2EShared: case.offset_provider[V2EShared]}) @pytest.mark.uses_unstructured_shift diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py index d882a88ec4..abad348c92 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_program.py @@ -60,7 +60,7 @@ def shift_by_one(in_field: cases.IFloatField) -> cases.IFloatField: # direct call to field operator # TODO(tehrengruber): slicing located fields not supported currently - # shift_by_one(in_field, out=out_field[:-1], offset_provider={"Ioff": IDim}) + # shift_by_one(in_field, out=out_field[:-1]) @gtx.program def shift_by_one_program(in_field: cases.IFloatField, out_field: cases.IFloatField): diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py index dd0cc6fb43..f0ce221851 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_reductions.py @@ -53,7 +53,7 @@ def testee(edge_f: cases.EField) -> cases.VField: inp = cases.allocate(unstructured_case, testee, "edge_f", strategy=strategy)() out = cases.allocate(unstructured_case, testee, cases.RETURN)() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() ref = np.max( inp.asnumpy()[v2e_table], axis=1, @@ -70,7 +70,7 @@ def minover(edge_f: cases.EField) -> cases.VField: out = min_over(edge_f(V2E), axis=V2EDim) return out - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, minover, @@ -100,7 +100,7 @@ def reduction_ek_field( "fop", [reduction_e_field, reduction_ek_field], ids=lambda fop: fop.__name__ ) def test_neighbor_sum(unstructured_case_3d, fop): - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() edge_f = cases.allocate(unstructured_case_3d, fop, "edge_f")() @@ -152,7 +152,7 @@ def fencil_op(edge_f: EKField) -> VKField: def fencil(edge_f: EKField, out: VKField): fencil_op(edge_f, out=out) - v2e_table = unstructured_case_3d.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case_3d.offset_provider[V2E].asnumpy() field = cases.allocate(unstructured_case_3d, fencil, "edge_f", sizes={KDim: 2})() out = cases.allocate(unstructured_case_3d, fencil_op, cases.RETURN, sizes={KDim: 1})() @@ -185,7 +185,7 @@ def reduce_expr(edge_f: cases.EField) -> cases.VField: def fencil(edge_f: cases.EField, out: cases.VField): reduce_expr(edge_f, out=out) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, fencil, @@ -207,7 +207,7 @@ def test_reduction_with_common_expression(unstructured_case): def testee(flux: cases.EField) -> cases.VField: return neighbor_sum(flux(V2E) + flux(V2E), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -223,7 +223,7 @@ def test_reduction_expression_with_where(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, inp(V2E), inp(V2E)), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -252,7 +252,7 @@ def test_reduction_expression_with_where_and_tuples(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(where(mask, (inp(V2E), inp(V2E)), (inp(V2E), inp(V2E)))[1], axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -281,7 +281,7 @@ def test_reduction_expression_with_where_and_scalar(unstructured_case): def testee(mask: cases.VBoolField, inp: cases.EField) -> cases.VField: return neighbor_sum(inp(V2E) + where(mask, inp(V2E), 1), axis=V2EDim) - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() mask = unstructured_case.as_field( [Vertex], np.random.choice(a=[False, True], size=unstructured_case.default_sizes[Vertex]) @@ -326,7 +326,7 @@ def testee(a: cases.VField) -> cases.EField: cases.verify_with_default_data( unstructured_case, testee, - ref=lambda a: a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], + ref=lambda a: a[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]], ) @@ -347,7 +347,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() neighbor_0_iter = iter(enumerate(e2v_table[:, 0])) edge_start = next(i for i, v in neighbor_0_iter if v >= ORIGIN) edge_stop = next(i for i, v in neighbor_0_iter if v < ORIGIN) @@ -392,16 +392,16 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_flat, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], ) cases.verify_with_default_data( unstructured_case, composed_shift_unstructured_intermediate_result, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], comparison=lambda inp, tmp: np.all(inp == tmp), ) @@ -409,8 +409,8 @@ def composed_shift_unstructured(inp: cases.VField) -> cases.CField: cases.verify_with_default_data( unstructured_case, composed_shift_unstructured, - ref=lambda inp: inp[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]][ - unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 0] + ref=lambda inp: inp[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]][ + unstructured_case.offset_provider[C2E].asnumpy()[:, 0] ], ) @@ -432,7 +432,7 @@ def testee(a: cases.VField) -> cases.EField: out = cases.allocate(unstructured_case, testee, cases.RETURN)() ORIGIN = 2 - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() neighbor_iter = iter(enumerate(e2v_table)) edge_start = next(i for i, v in neighbor_iter if all(v >= ORIGIN)) edge_stop = next(i for i, v in neighbor_iter if any(v < ORIGIN)) @@ -453,12 +453,11 @@ def testee(a: cases.VField) -> cases.VField: unstructured_case, testee, ref=lambda a: np.sum( - np.sum(a[unstructured_case.offset_provider[E2VDim.tag].asnumpy()], axis=1, initial=0)[ - unstructured_case.offset_provider[V2EDim.tag].asnumpy() + np.sum(a[unstructured_case.offset_provider[E2V].asnumpy()], axis=1, initial=0)[ + unstructured_case.offset_provider[V2E].asnumpy() ], axis=1, - where=unstructured_case.offset_provider[V2EDim.tag].asnumpy() - != common._DEFAULT_SKIP_VALUE, + where=unstructured_case.offset_provider[V2E].asnumpy() != common._DEFAULT_SKIP_VALUE, ), comparison=lambda a, tmp_2: np.all(a == tmp_2), ) @@ -479,8 +478,8 @@ def testee(inp: cases.EField) -> cases.EField: unstructured_case, testee, ref=lambda inp: np.sum( - np.sum(inp[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1)[ - unstructured_case.offset_provider[E2VDim.tag].asnumpy() + np.sum(inp[unstructured_case.offset_provider[V2E].asnumpy()], axis=1)[ + unstructured_case.offset_provider[E2V].asnumpy() ], axis=1, ), @@ -497,7 +496,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: tmp = red(E2V[0]) return tmp - v2e = unstructured_case.offset_provider[V2EDim.tag] + v2e = unstructured_case.offset_provider[V2E] cases.verify_with_default_data( unstructured_case, reduce_tuple_element, @@ -506,7 +505,7 @@ def reduce_tuple_element(e: cases.EField, v: cases.VField) -> cases.EField: axis=1, initial=0, where=v2e.asnumpy() != common._DEFAULT_SKIP_VALUE, - )[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 0]], + )[unstructured_case.offset_provider[E2V].asnumpy()[:, 0]], ) @@ -518,7 +517,7 @@ def testee(a: cases.EField, b: cases.EField) -> cases.VField: tmp = neighbor_sum(b(V2E) if 2 < 3 else a(V2E), axis=V2EDim) return tmp - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify_with_default_data( unstructured_case, testee, @@ -540,7 +539,7 @@ def testee(inp: gtx.Field[[Edge], int32]) -> gtx.Field[[Vertex], int32]: inp = cases.allocate(unstructured_case, testee, "inp")() - v2e_table = unstructured_case.offset_provider[V2EDim.tag].asnumpy() + v2e_table = unstructured_case.offset_provider[V2E].asnumpy() cases.verify( unstructured_case, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py index ac348266f3..15a94f2e36 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_staggered.py @@ -26,7 +26,7 @@ unstructured_case, unstructured_case_3d, ) -from next_tests.integration_tests.cases_utils import Koff, exec_alloc_descriptor, mesh_descriptor +from next_tests.integration_tests.cases_utils import exec_alloc_descriptor, mesh_descriptor @pytest.mark.uses_cartesian_shift @@ -156,7 +156,7 @@ def test_cartesian_half_shift_as_offset(cartesian_case): def testee( a: gtx.Field[[IDim, KHalfDim], np.int32], offset_field: cases.IKField ) -> cases.IKField: - return a(KDim - 0.5)(as_offset(Koff, offset_field)) + return a(KDim - 0.5)(as_offset(KDim, offset_field)) ksize = cartesian_case.default_sizes[KDim] a = cases.allocate(cartesian_case, testee, "a", sizes={KHalfDim: ksize + 1})() @@ -180,7 +180,7 @@ def testee( a: gtx.Field[[IDim, KHalfDim], np.int32], offset_field: cases.IKField ) -> cases.IKField: b = a + 1 - return b(KDim - 0.5)(as_offset(Koff, offset_field)) + return b(KDim - 0.5)(as_offset(KDim, offset_field)) ksize = cartesian_case.default_sizes[KDim] a = cases.allocate(cartesian_case, testee, "a", sizes={KHalfDim: ksize + 1})() @@ -208,7 +208,7 @@ def testee( a: gtx.Field[[Vertex, KHalfDim], np.int32], offset_field: gtx.Field[[Edge, KDim], np.int32], ) -> gtx.Field[[Edge, KDim], np.int32]: - return a(E2V[0])(KDim - 0.5)(as_offset(Koff, offset_field)) + return a(E2V[0])(KDim - 0.5)(as_offset(KDim, offset_field)) nvertices = unstructured_case_3d.default_sizes[Vertex] ksize = unstructured_case_3d.default_sizes[KDim] @@ -223,7 +223,7 @@ def testee( )() out = cases.allocate(unstructured_case_3d, testee, cases.RETURN)() - e2v_table = unstructured_case_3d.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case_3d.offset_provider[E2V].asnumpy() cases.verify( unstructured_case_3d, testee, diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py index 3db69adf0e..1b75ef54e5 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py @@ -82,9 +82,7 @@ def test_verification(testee, exec_alloc_descriptor, mesh_descriptor): a = cases.allocate(unstructured_case, testee, "a")() out = cases.allocate(unstructured_case, testee, "out")() - first_nbs, second_nbs = ( - mesh_descriptor.offset_provider[E2VDim.tag].asnumpy()[:, i] for i in [0, 1] - ) + first_nbs, second_nbs = (mesh_descriptor.offset_provider[E2V].asnumpy()[:, i] for i in [0, 1]) ref = (a.ndarray * 2)[first_nbs] + (a.ndarray * 2)[second_nbs] cases.verify( @@ -104,7 +102,7 @@ def test_temporary_symbols(testee, mesh_descriptor): gtir_with_tmp = apply_common_transforms( testee.gtir, extract_temporaries=True, - offset_provider=mesh_descriptor.offset_provider, + offset_provider=common.as_tag_keyed_offset_provider(mesh_descriptor.offset_provider), ) params = ["num_vertices", "num_edges", "num_cells"] diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py index c55a145314..ed91a34e9e 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_tuples.py @@ -164,8 +164,8 @@ def testee(a: cases.EField, b: cases.EField) -> tuple[cases.VField, cases.VField unstructured_case, testee, ref=lambda a, b: [ - np.sum(a[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), - np.sum(b[unstructured_case.offset_provider[V2EDim.tag].asnumpy()], axis=1), + np.sum(a[unstructured_case.offset_provider[V2E].asnumpy()], axis=1), + np.sum(b[unstructured_case.offset_provider[V2E].asnumpy()], axis=1), ], comparison=lambda a, tmp: (np.all(a[0] == tmp[0]), np.all(a[1] == tmp[1])), ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py index 6a160cefc3..2e4f3d8b78 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_type_conversion.py @@ -51,7 +51,7 @@ def testee(a: gtx.Field[[Vertex], np.float64]) -> gtx.Field[[Edge], int64]: tmp = astype(a(E2V), int64) return neighbor_sum(tmp, axis=E2VDim) - e2v_table = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + e2v_table = unstructured_case.offset_provider[E2V].asnumpy() cases.verify_with_default_data( unstructured_case, diff --git a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py index 01e4ed92bb..eabf67327d 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/ffront_tests/test_multiple_output_domains.py @@ -546,7 +546,7 @@ def test_program_unstructured(unstructured_case): unstructured_case.default_sizes[Cell], unstructured_case.default_sizes[Edge], inout=(out_a_shifted, out_a), - ref=((a.ndarray)[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]], a), + ref=((a.ndarray)[unstructured_case.offset_provider[C2E].asnumpy()[:, 1]], a), ) @@ -600,7 +600,7 @@ def test_program_temporary(unstructured_case): extend={Cell: (-restrict_cell[0], restrict_cell[1])}, )() - e2v = (a.ndarray)[unstructured_case.offset_provider[E2VDim.tag].asnumpy()[:, 1]] + e2v = (a.ndarray)[unstructured_case.offset_provider[E2V].asnumpy()[:, 1]] cases.verify( unstructured_case, prog_temporary, @@ -616,7 +616,7 @@ def test_program_temporary(unstructured_case): inout=(out_edge, out_cell), ref=( e2v[restrict_edge[0] : edge_size + restrict_edge[1]], - e2v[unstructured_case.offset_provider[C2EDim.tag].asnumpy()[:, 1]][ + e2v[unstructured_case.offset_provider[C2E].asnumpy()[:, 1]][ restrict_cell[0] : cell_size + restrict_cell[1] ], ), diff --git a/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py b/tests/next_tests/integration_tests/multi_feature_tests/fvm_nabla_setup.py index 99d8085dd0..4c9edb0e95 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 0029) a same-named declaration # here would be a different dimension from the one `toy_connectivity` declares, where the old # `Dimension("...")` values compared equal -- and tests mix objects from both modules. -from next_tests.toy_connectivity import E2VDim, Edge, V2EDim, Vertex - - -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 b7c81d55fd..91f0e9b2a9 100644 --- a/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py +++ b/tests/next_tests/regression_tests/ffront_tests/test_offset_dimensions_names.py @@ -7,19 +7,13 @@ # SPDX-License-Identifier: BSD-3-Clause """ -Regression tests for the four independently authored names of one connectivity. +Regression tests for the names under which one connectivity is used. -Using a single connectivity requires four strings to agree, none of which is -checked against the others at declaration time: - - N1 the `FieldOffset` tag `FieldOffset("V2E", ...)` - N2 the Python variable it is bound to `V2E = FieldOffset(...)` - N3 the local dimension's name `Dimension("V2E", kind=LOCAL)` - N4 the offset-provider key `offset_provider={"V2E": ...}` - -The `V2EDim = Dimension("V2E")` convention makes all four equal, which hides -which one each execution path actually uses. These tests break the convention -deliberately, one name at a time, so the real requirement is visible. +With `FieldOffset`, using a connectivity required four independently authored strings to +agree -- the offset tag, the Python variable it was bound to, the local dimension's name and +the offset-provider key -- and each execution path silently depended on a different subset of +them. A `NeighborConnectivity` declaration produces all of them (ADR 0030), so what is left to +pin is that the *Python* name a declaration is reached through does not matter. """ import numpy as np @@ -42,26 +36,22 @@ class V(gtx.DimensionIndex): ... class E(gtx.DimensionIndex): ... -#: N1 == N3 == N4, but N2 differs: the tag is `TaggedOffDim.tag`, the variable is `off_a`. -class TaggedOffDim(gtx.LocalDimensionIndex): ... - - -off_a = gtx.FieldOffset(TaggedOffDim.tag, source=E, target=(V, TaggedOffDim)) +class V2E(gtx.NeighborConnectivity[V, E]): + class Local(gtx.LocalDimensionIndex): ... -#: N1 == N2 == N4, but N3 differs: the local dimension is `Neigh`, the tag is `OffB`. -class Neigh(gtx.LocalDimensionIndex): ... +#: The declaration, reached through a different Python name. +off_a = V2E +#: Its local dimension, likewise. +Neigh = V2E.Local -OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) - - -def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimension) -> cases.Case: - """A `Case` binding the same table under each of `tags`.""" +@pytest.fixture +def case(exec_alloc_descriptor) -> cases.Case: mesh = cases_utils.simple_mesh(exec_alloc_descriptor.allocator) # NOTE: `.asnumpy()`, not `.ndarray`: under a GPU allocator the latter is a device # array, and `simple_mesh` builds the table from NumPy anyway. - v2e_arr = mesh.offset_provider[cases_utils.V2EDim.tag].asnumpy() + v2e_arr = mesh.offset_provider[cases_utils.V2E].asnumpy() return cases.Case( ( None @@ -69,14 +59,13 @@ def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimens else exec_alloc_descriptor ), offset_provider={ - tag: constructors.as_connectivity( - domain={V: v2e_arr.shape[0], local_dim: v2e_arr.shape[1]}, + off_a: constructors.as_connectivity( + domain={V: v2e_arr.shape[0], Neigh: v2e_arr.shape[1]}, codomain=E, data=v2e_arr, skip_value=None, allocator=exec_alloc_descriptor.allocator, ) - for tag in tags }, default_sizes={V: mesh.num_vertices, E: mesh.num_edges}, grid_type=common.GridType.UNSTRUCTURED, @@ -84,82 +73,29 @@ def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimens ) -@pytest.fixture -def case_tag_vs_variable_name(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, (TaggedOffDim.tag,), TaggedOffDim) - - -@pytest.fixture -def case_tag_vs_local_dim(exec_alloc_descriptor): - # NOTE: only the offset's table: a reduction over `Neigh` finds it as the table over `Neigh`, - # as for a connectivity sharing another one's local dimension. - return _case(exec_alloc_descriptor, ("OffB",), Neigh) - - -def _neighbor_table(case: cases.Case, tag: str) -> np.ndarray: - return case.offset_provider[tag].asnumpy() - +def _neighbor_table(case: cases.Case) -> np.ndarray: + return case.offset_provider[V2E].asnumpy() -# --- N2: the tag differs from the Python variable name ---------------------------- -# Lowering used to emit the *variable* name as the IR shift tag, so embedded and -# compiled execution of the same program needed different provider keys. - -def test_shift_tag_differs_from_variable_name(case_tag_vs_variable_name): +def test_shift_through_an_alias(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: return a(off_a[1]) - cases.verify_with_default_data( - case_tag_vs_variable_name, - foo, - lambda a: a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)[:, 1]], - ) - - -def test_reduction_tag_differs_from_variable_name(case_tag_vs_variable_name): - @gtx.field_operator - def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return neighbor_sum(a(off_a), axis=TaggedOffDim) - - cases.verify_with_default_data( - case_tag_vs_variable_name, - foo, - lambda a: np.sum(a[_neighbor_table(case_tag_vs_variable_name, TaggedOffDim.tag)], axis=1), - ) - - -# --- N3: the tag differs from the local dimension's name -------------------------- -# The shape of a connectivity sharing another one's local dimension. - + cases.verify_with_default_data(case, foo, lambda a: a[_neighbor_table(case)[:, 1]]) -def test_shift_tag_differs_from_local_dim_name(case_tag_vs_local_dim): - """ - Ensure a shift works with an offset tag that differs from the local dimension's name. - - If the local dimension of the `NeighborTableType` did not match the `FieldOffset` value, - gtfn would silently ignore the neighbor index, see - https://github.com/GridTools/gridtools/pull/1814. - """ +def test_reduction_through_an_alias(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return a(OffB[1]) + return neighbor_sum(a(off_a), axis=Neigh) - cases.verify_with_default_data( - case_tag_vs_local_dim, - foo, - lambda a: a[_neighbor_table(case_tag_vs_local_dim, "OffB")[:, 1]], - ) + cases.verify_with_default_data(case, foo, lambda a: np.sum(a[_neighbor_table(case)], axis=1)) -def test_reduction_tag_differs_from_local_dim_name(case_tag_vs_local_dim): +def test_reduction_over_the_nested_name(case): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: - return neighbor_sum(a(OffB), axis=Neigh) + return neighbor_sum(a(V2E), axis=off_a.Local) - cases.verify_with_default_data( - case_tag_vs_local_dim, - foo, - lambda a: np.sum(a[_neighbor_table(case_tag_vs_local_dim, "OffB")], axis=1), - ) + cases.verify_with_default_data(case, foo, lambda a: np.sum(a[_neighbor_table(case)], axis=1)) diff --git a/tests/next_tests/toy_connectivity.py b/tests/next_tests/toy_connectivity.py index 531ee0867f..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.LocalDimensionIndex): ... +class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class E2VDim(gtx.LocalDimensionIndex): ... +class E2V(gtx.NeighborConnectivity[Edge, Vertex]): + class Local(gtx.LocalDimensionIndex): ... -class C2EDim(gtx.LocalDimensionIndex): ... +class C2E(gtx.NeighborConnectivity[Cell, Edge]): + class Local(gtx.LocalDimensionIndex): ... -class V2VDim(gtx.LocalDimensionIndex): ... +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 1cbb1818c2..2149d9f8ac 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 @@ -22,6 +22,7 @@ DimensionIndex, LocalDimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Field, DimensionIndex, @@ -828,7 +829,6 @@ def test_premap_disjoint_inverse_image_raises(): def test_as_offset_1d(): # Dynamic per-point shift along I: out[i] == f[i + off[i]], full domain when all shifts in-bounds. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -836,7 +836,7 @@ def test_as_offset_1d(): off_arr = np.asarray([1, 0, -1, 0, 1, 0, -1, 0, 1, 0], dtype=int) off = common._field(off_arr, domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),))) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) assert np.all(result.ndarray == f.ndarray[np.arange(10) + off_arr]) @@ -844,7 +844,6 @@ def test_as_offset_1d(): def test_as_offset_narrow_offset_dtype_no_wrap(): # An int8 offset field over a domain larger than 128 must not wrap into the index table. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) N = 200 f = common._field( @@ -853,7 +852,7 @@ def test_as_offset_narrow_offset_dtype_no_wrap(): off_arr = np.zeros(N, dtype=np.int8) off = common._field(off_arr, domain=common.Domain(dims=(I,), ranges=(UnitRange(0, N),))) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(0, N),)) assert np.all(result.ndarray == f.ndarray) @@ -861,7 +860,6 @@ def test_as_offset_narrow_offset_dtype_no_wrap(): def test_as_offset_2d_shift_one_keep_other(): # Shift along I by a per-(i, j) offset, leave J: out[i, j] == f[i + off[i, j], j]. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) NI, NJ = 4, 3 dom = common.Domain(dims=(I, J), ranges=(UnitRange(0, NI), UnitRange(0, NJ))) @@ -869,7 +867,7 @@ def test_as_offset_2d_shift_one_keep_other(): off_arr = np.asarray([[1, 1, 0], [0, 0, 1], [1, -1, 0], [-1, 0, -1]], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == dom i = np.arange(NI)[:, None] @@ -879,7 +877,6 @@ def test_as_offset_2d_shift_one_keep_other(): def test_as_offset_boundary_narrows_domain(): # A uniform out-of-bounds shift narrows the result to the contiguous in-range sub-domain. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -888,7 +885,7 @@ def test_as_offset_boundary_narrows_domain(): np.full(10, -1, dtype=int), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) ) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I,), ranges=(UnitRange(1, 10),)) assert np.all(result.ndarray == f.ndarray[0:9]) # out[i] == f[i - 1] @@ -896,7 +893,6 @@ def test_as_offset_boundary_narrows_domain(): def test_as_offset_scattered_oob_raises(): # An out-of-bounds shift in the interior cannot yield a contiguous domain. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -907,12 +903,11 @@ def test_as_offset_scattered_oob_raises(): ) with pytest.raises(ValueError, match="non-contiguous"): - f.premap(as_offset(Ioff, off)) + f.premap(as_offset(I, off)) def test_as_offset_introduces_dimension(): # `off` carries a dim the field lacks: the result gains it, out[i, j] == f[i + off[i, j]]. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) f = common._field( np.arange(10).astype(float), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 10),)) @@ -922,7 +917,7 @@ def test_as_offset_introduces_dimension(): off_arr, domain=common.Domain(dims=(I, J), ranges=(UnitRange(0, 10), UnitRange(0, 3))) ) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == common.Domain(dims=(I, J), ranges=(UnitRange(0, 10), UnitRange(0, 3))) assert np.all(result.ndarray == f.ndarray[np.arange(10)[:, None] + off_arr]) @@ -930,14 +925,13 @@ def test_as_offset_introduces_dimension(): def test_as_offset_nonzero_origin(): # Field and offset over a domain that does not start at 0: indices must be shifted by the domain start. - Ioff = fbuiltins.FieldOffset("Ioff", source=I, target=(I,)) dom = common.Domain(dims=(I,), ranges=(UnitRange(2, 12),)) f = common._field(np.arange(10).astype(float), domain=dom) off_arr = np.asarray([1, 0, -1, 0, 1, 0, -1, 0, 1, 0], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Ioff, off)) + result = f.premap(as_offset(I, off)) assert result.domain == dom assert np.all(result.ndarray == f.ndarray[np.arange(10) + off_arr]) @@ -945,7 +939,6 @@ def test_as_offset_nonzero_origin(): def test_as_offset_2d_shift_second_axis(): # Shift along J (the non-leading axis) by a per-(i, j) offset, leave I: out[i, j] == f[i, j + off[i, j]]. - Joff = fbuiltins.FieldOffset("Joff", source=J, target=(J,)) NI, NJ = 3, 4 dom = common.Domain(dims=(I, J), ranges=(UnitRange(0, NI), UnitRange(0, NJ))) @@ -953,7 +946,7 @@ def test_as_offset_2d_shift_second_axis(): off_arr = np.asarray([[1, 1, 0, -1], [0, 0, 1, -1], [1, -1, 0, 0]], dtype=int) off = common._field(off_arr, domain=dom) - result = f.premap(as_offset(Joff, off)) + result = f.premap(as_offset(J, off)) assert result.domain == dom i = np.arange(NI)[:, None] @@ -961,25 +954,14 @@ def test_as_offset_2d_shift_second_axis(): assert np.all(result.ndarray == f.ndarray[i, j + off_arr]) -def test_as_offset_non_cartesian_offset_raises(): - # `as_offset` only supports Cartesian (self-shift) offsets: single target equal to source. - - off_I = common._field( - np.zeros(3, dtype=int), domain=common.Domain(dims=(I,), ranges=(UnitRange(0, 3),)) - ) +@pytest.mark.parametrize("dim", [V2EDim, Vertex]) +def test_as_offset_off_an_axis_raises(dim): + # `as_offset` shifts along a Cartesian axis: not a local dimension, nor a mesh location. off_V = common._field( np.zeros(3, dtype=int), domain=common.Domain(dims=(Vertex,), ranges=(UnitRange(0, 3),)) ) - - # 2-element target (neighbor offset) - V2E = fbuiltins.FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) - with pytest.raises(ValueError, match="Cartesian"): - as_offset(V2E, off_V) - - # 1-element target but source != target[0] (cross-dim) - IfromJ = fbuiltins.FieldOffset("IfromJ", source=I, target=(J,)) - with pytest.raises(ValueError, match="Cartesian"): - as_offset(IfromJ, off_I) + with pytest.raises(TypeError, match="shifts along a Cartesian axis"): + as_offset(dim, off_V) @pytest.mark.parametrize( diff --git a/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py b/tests/next_tests/unit_tests/ffront_tests/test_decorator_domain_deduction.py index 7c8c026f58..fd448fd1a6 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 @@ -24,30 +26,25 @@ class Dim(gtx.DimensionIndex): ... class LocalDim(gtx.LocalDimensionIndex): ... -CartesianOffset = gtx.FieldOffset("CartesianOffset", source=Dim, target=(Dim,)) -UnstructuredOffset = gtx.FieldOffset(LocalDim.tag, source=Dim, target=(Dim, LocalDim)) +class UnstructuredOffset(gtx.NeighborConnectivity[Dim, Dim]): + Local: typing.TypeAlias = LocalDim def test_domain_deduction_cartesian(): - assert _deduce_grid_type(None, {CartesianOffset}) == gtx.GridType.CARTESIAN assert _deduce_grid_type(None, {Dim}) == gtx.GridType.CARTESIAN + assert _deduce_grid_type(None, {HDim, VDim}) == gtx.GridType.CARTESIAN def test_domain_deduction_unstructured(): assert _deduce_grid_type(None, {UnstructuredOffset}) == gtx.GridType.UNSTRUCTURED assert _deduce_grid_type(None, {LocalDim}) == gtx.GridType.UNSTRUCTURED - # source and target share `.value` but differ in `.kind` -> not Cartesian - CrossKindOffset = gtx.FieldOffset("CrossKind", source=HDim, target=(VDim,)) - assert _deduce_grid_type(None, {CrossKindOffset}) == gtx.GridType.UNSTRUCTURED - # LOCAL self-loop is unstructured - LocalSelfOffset = gtx.FieldOffset("LocalSelf", source=LocalDim, target=(LocalDim,)) - assert _deduce_grid_type(None, {LocalSelfOffset}) == gtx.GridType.UNSTRUCTURED def test_domain_complies_with_request_cartesian(): - assert _deduce_grid_type(gtx.GridType.CARTESIAN, {CartesianOffset}) == gtx.GridType.CARTESIAN - with pytest.raises(ValueError, match="unstructured.*FieldOffset.*found"): + assert _deduce_grid_type(gtx.GridType.CARTESIAN, {Dim}) == gtx.GridType.CARTESIAN + with pytest.raises(ValueError, match="NeighborConnectivity.*local dimension was found"): _deduce_grid_type(gtx.GridType.CARTESIAN, {UnstructuredOffset}) + with pytest.raises(ValueError, match="NeighborConnectivity.*local dimension was found"): _deduce_grid_type(gtx.GridType.CARTESIAN, {LocalDim}) @@ -57,6 +54,4 @@ def test_domain_complies_with_request_unstructured(): == gtx.GridType.UNSTRUCTURED ) # unstructured is ok, even if we don't have unstructured offsets - assert ( - _deduce_grid_type(gtx.GridType.UNSTRUCTURED, {CartesianOffset}) == gtx.GridType.UNSTRUCTURED - ) + assert _deduce_grid_type(gtx.GridType.UNSTRUCTURED, {Dim}) == gtx.GridType.UNSTRUCTURED diff --git a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py index 3a272b156e..8cdb550128 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_diagnostic_messages.py @@ -30,8 +30,6 @@ class IDim(gtx.CartesianAxisIndex): ... -IOff = gtx.FieldOffset("Ioff", source=IDim, target=(IDim,)) - # A PEP 695 alias whose value raises when it is evaluated, standing in for the # common case of a typo'd dtype ('np.foat64') inside an alias definition. _empty_module = types.ModuleType("_empty_module") @@ -361,20 +359,6 @@ def broken(a: BrokenFieldAlias) -> gtx.Field[[IDim], float64]: assert re.search(r"\| +\^{19}(?!\^)", str(err)), str(err) -def test_unindexed_cartesian_offset_names_the_offset_as_written(): - # The tag of 'IOff' is 'Ioff'; the message has to quote what the user wrote. - def unindexed(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: - return a(IOff) - - err = parse_error(unindexed) - - assert err.message == "Cannot shift by the Cartesian offset 'IOff' without an index." - assert err.hints == ["Give the displacement, e.g. 'IOff[1]'."] - rendered = str(err) - assert "return a(IOff)" in rendered - assert re.search(r"\| +\^{4}(?!\^)", rendered), rendered - - def test_indexed_dimension_shift_is_rejected_with_a_hint(): def indexed_shift(a: gtx.Field[[IDim], float64]) -> gtx.Field[[IDim], float64]: return a((IDim + 1)[0]) diff --git a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py index 37c650d998..81bc782469 100644 --- a/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py +++ b/tests/next_tests/unit_tests/ffront_tests/test_foast_to_gtir.py @@ -49,24 +49,19 @@ class Edge(gtx.DimensionIndex): ... class Vertex(gtx.DimensionIndex): ... -class V2EDim(gtx.LocalDimensionIndex): ... +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.CartesianAxisIndex): ... -TOff = gtx.FieldOffset("TDim", source=TDim, target=(TDim,)) - - -#: An offset whose tag differs from the name of the Python variable it is bound to, and -#: from the name of its local dimension. Lowering must emit the *tag*. -class RenamedV2EDim(gtx.LocalDimensionIndex): ... - - -renamed_v2e = gtx.FieldOffset("RenamedTag", source=Edge, target=(Vertex, RenamedV2EDim)) +#: A connectivity reached through a name other than its declaration's. Lowering must emit the +#: local dimension's tag, not the variable name. +renamed_v2e = V2E class UDim(gtx.CartesianAxisIndex): ... @@ -166,7 +161,7 @@ def foo_float(inp: gtx.Field[[TDim], float64]): def test_as_offset(): def foo(inp: gtx.Field[[TDim], float64], offset: gtx.Field[[TDim], int]): - return inp(as_offset(TOff, offset)) + return inp(as_offset(TDim, offset)) parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) @@ -802,7 +797,7 @@ def foo(edge_f: gtx.Field[gtx.Dims[Edge], float64]): def test_unstructured_shift_lowering_emits_offset_tag_not_variable_name(): - """The IR shift tag is the offset's tag, not the variable the offset is bound to.""" + """The IR shift tag is the local dimension's tag, not the variable it is reached through.""" def foo(edge_f: gtx.Field[[Edge], float64]): return edge_f(renamed_v2e[1]) @@ -810,7 +805,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) - reference = im.as_fieldop(im.lambda_("__it")(im.deref(im.shift("RenamedTag", 1)("__it"))))( + reference = im.as_fieldop(im.lambda_("__it")(im.deref(im.shift(V2EDim.tag, 1)("__it"))))( "edge_f" ) @@ -824,7 +819,7 @@ def foo(edge_f: gtx.Field[[Edge], float64]): parsed = FieldOperatorParser.apply_to_function(foo) lowered = FieldOperatorLowering.apply(parsed) - reference = im.as_fieldop_neighbors("RenamedTag", "edge_f") + reference = im.as_fieldop_neighbors(V2EDim.tag, "edge_f") assert lowered.expr == reference 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 2856c31149..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 @@ -30,10 +30,11 @@ class Cell(DimensionIndex): ... class Edge(DimensionIndex): ... -class C2EDim(LocalDimensionIndex): ... +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 afa3bba76e..1f15e57105 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 @@ -20,8 +20,9 @@ DimensionIndex, LocalDimensionIndex, DimensionKind, + LocalDimensionIndex, + NeighborConnectivity, Field, - FieldOffset, astype, broadcast, errors, @@ -52,7 +53,11 @@ class X(CartesianAxisIndex): ... class Y(CartesianAxisIndex): ... -class Y2XDim(LocalDimensionIndex): ... +class Y2X(NeighborConnectivity[Y, X]): + class Local(LocalDimensionIndex): ... + + +Y2XDim = Y2X.Local class K(CartesianAxisIndex, kind=DimensionKind.VERTICAL): ... @@ -73,7 +78,11 @@ class Vertex(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... class Edge(DimensionIndex, kind=DimensionKind.HORIZONTAL): ... -class V2EDim(LocalDimensionIndex): ... +class V2E(NeighborConnectivity[Vertex, Edge]): + class Local(LocalDimensionIndex): ... + + +V2EDim = V2E.Local class IDim(CartesianAxisIndex): ... @@ -288,7 +297,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 @@ -549,41 +557,39 @@ def return_undefined(): def test_as_offset_dim(): - Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) - def as_offset_dim(a: Field[[ADim, BDim], float], b: Field[[ADim], int]): - return a(as_offset(Boff, b)) + return a(as_offset(BDim, b)) with pytest.raises(errors.DSLError, match=f"not in list of offset field dimensions"): _ = FieldOperatorParser.apply_to_function(as_offset_dim) def test_as_offset_dtype(): - Boff = FieldOffset("Boff", source=BDim, target=(BDim,)) - def as_offset_dtype(a: Field[[ADim, BDim], float], b: Field[[BDim], float]): - return a(as_offset(Boff, b)) + return a(as_offset(BDim, b)) with pytest.raises(errors.DSLError, match=f"expected integer for offset field dtype"): _ = FieldOperatorParser.apply_to_function(as_offset_dtype) -def test_as_offset_non_cartesian(): - V2E = FieldOffset(V2EDim.tag, source=Edge, target=(Vertex, V2EDim)) - +def test_as_offset_non_dimension(): def as_offset_neighbor(a: Field[[Edge], float], b: Field[[Edge], int]): return a(as_offset(V2E, b)) - with pytest.raises(errors.DSLError, match="Cartesian"): + with pytest.raises(errors.DSLError, match="Expected 1st argument to be of type"): _ = FieldOperatorParser.apply_to_function(as_offset_neighbor) - IfromJ = FieldOffset("IfromJ", source=IDim, target=(JDim,)) + def as_offset_local_dim(a: Field[[Vertex, V2EDim], float], b: Field[[Vertex, V2EDim], int]): + return a(as_offset(V2EDim, b)) + + with pytest.raises(errors.DSLError, match="shifts along a Cartesian axis"): + _ = FieldOperatorParser.apply_to_function(as_offset_local_dim) - def as_offset_cross_dim(a: Field[[IDim], float], b: Field[[IDim], int]): - return a(as_offset(IfromJ, b)) + def as_offset_mesh_location(a: Field[[Edge], float], b: Field[[Edge], int]): + return a(as_offset(Edge, b)) - with pytest.raises(errors.DSLError, match="Cartesian"): - _ = FieldOperatorParser.apply_to_function(as_offset_cross_dim) + with pytest.raises(errors.DSLError, match="shifts along a Cartesian axis"): + _ = FieldOperatorParser.apply_to_function(as_offset_mesh_location) @pytest.mark.parametrize("offset", [1, -1, 0.5]) 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 33f6f71f43..842e84fc4e 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_field_with_list.py @@ -30,10 +30,11 @@ class E(gtx.DimensionIndex): ... class V(gtx.DimensionIndex): ... -class E2VDim(gtx.LocalDimensionIndex): ... +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/test_inline_dynamic_shifts.py b/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py index 4c7b568f1e..0581637946 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_inline_dynamic_shifts.py @@ -27,9 +27,7 @@ def test_inline_dynamic_shift_as_fieldop_arg(uids): im.lambda_("inp", "offset_field")(im.deref(im.shift(IOff, im.deref("offset_field"))("inp"))) )("inp", "offset_field") - actual = inline_dynamic_shifts.InlineDynamicShifts.apply( - testee, offset_provider_type={}, uids=uids - ) + actual = inline_dynamic_shifts.InlineDynamicShifts.apply(testee, table_types={}, uids=uids) assert actual == expected @@ -46,9 +44,7 @@ def test_inline_dynamic_shift_nested_as_fieldop_args(uids): ) )("inp", "offset_field") - actual = inline_dynamic_shifts.InlineDynamicShifts.apply( - testee, offset_provider_type={}, uids=uids - ) + actual = inline_dynamic_shifts.InlineDynamicShifts.apply(testee, table_types={}, uids=uids) assert actual == expected @@ -63,7 +59,5 @@ def test_inline_dynamic_shift_let_var(uids): im.lambda_("inp", "offset_field")(im.deref(im.shift(IOff, im.deref("offset_field"))("inp"))) )("inp", "offset_field") - actual = inline_dynamic_shifts.InlineDynamicShifts.apply( - testee, offset_provider_type={}, uids=uids - ) + actual = inline_dynamic_shifts.InlineDynamicShifts.apply(testee, table_types={}, uids=uids) assert actual == expected diff --git a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py index 4f580b110c..cc17127979 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_type_inference.py @@ -309,11 +309,11 @@ def expression_test_cases(): @pytest.mark.parametrize("test_case", expression_test_cases()) def test_expression_type(test_case): mesh = simple_mesh(None) - offset_provider_type = mesh.offset_provider_type + table_types = mesh.table_types testee, expected_type = test_case result = itir_type_inference.infer( - testee, offset_provider_type=offset_provider_type, allow_undeclared_symbols=True + testee, table_types=table_types, allow_undeclared_symbols=True ) assert result.type == expected_type @@ -324,17 +324,17 @@ def test_expression_type(test_case): ) def test_expression_type_as_fieldop_no_domain(test_case): mesh = simple_mesh(None) - offset_provider_type = mesh.offset_provider_type + table_types = mesh.table_types testee_with_domain, expected_type = test_case result_with_domain = itir_type_inference.infer( - testee_with_domain, offset_provider_type=offset_provider_type, allow_undeclared_symbols=True + testee_with_domain, table_types=table_types, allow_undeclared_symbols=True ) # testee stays as is, but we remove the domain testee_without_domain = im.as_fieldop(testee_with_domain.fun.args[0])(*testee_with_domain.args) result_without_domain = itir_type_inference.infer( testee_without_domain, - offset_provider_type=offset_provider_type, + table_types=table_types, allow_undeclared_symbols=True, ) assert result_with_domain.type == result_without_domain.type == expected_type @@ -344,9 +344,7 @@ def test_adhoc_polymorphism(): func = im.lambda_("a")(im.lambda_("b")(im.make_tuple("a", "b"))) testee = im.call(im.call(func)(im.ref("a_", bool_type)))(im.ref("b_", int_type)) - result = itir_type_inference.infer( - testee, offset_provider_type={}, allow_undeclared_symbols=True - ) + result = itir_type_inference.infer(testee, table_types={}, allow_undeclared_symbols=True) assert result.type == ts.TupleType(types=[bool_type, int_type]) @@ -355,9 +353,7 @@ def test_binary_lambda(): func = im.lambda_("a", "b")(im.make_tuple("a", "b")) testee = im.call(func)(im.ref("a_", bool_type), im.ref("b_", int_type)) - result = itir_type_inference.infer( - testee, offset_provider_type={}, allow_undeclared_symbols=True - ) + result = itir_type_inference.infer(testee, table_types={}, allow_undeclared_symbols=True) expected_type = ts.TupleType(types=[bool_type, int_type]) assert result.type == expected_type @@ -373,7 +369,7 @@ def test_binary_lambda(): def test_aliased_function(): testee = im.let("f", im.lambda_("x")("x"))(im.call("f")(1)) - result = itir_type_inference.infer(testee, offset_provider_type={}) + result = itir_type_inference.infer(testee, table_types={}) assert result.args[0].type == ts.FunctionType( pos_only_args=[int_type], pos_or_kw_args={}, kw_only_args={}, returns=int_type @@ -389,7 +385,7 @@ def test_late_offset_axis(): testee = im.call(func)(im.ensure_offset(V2EDim.tag)) result = itir_type_inference.infer( - testee, offset_provider_type=mesh.offset_provider_type, allow_undeclared_symbols=True + testee, table_types=mesh.table_types, allow_undeclared_symbols=True ) assert result.type == it_on_e_of_e_type @@ -399,9 +395,7 @@ def test_cast_first_arg_inference(): # easy to forget inferring the types of the first argument and its children. Simply check # if the first argument has a type inferred correctly here. testee = im.cast_(im.plus(im.literal_from_value(1), im.literal_from_value(2)), "float64") - result = itir_type_inference.infer( - testee, offset_provider_type={}, allow_undeclared_symbols=True - ) + result = itir_type_inference.infer(testee, table_types={}, allow_undeclared_symbols=True) assert result.args[0].type == int_type assert result.type == float64_type @@ -426,7 +420,7 @@ def test_cartesian_fencil_definition(): ], ) - result = itir_type_inference.infer(testee, offset_provider_type={}) + result = itir_type_inference.infer(testee, table_types={}) program_type = it_ts.ProgramType(params={"inp": float_i_field, "out": float_i_field}) assert result.type == program_type @@ -459,7 +453,7 @@ def test_unstructured_fencil_definition(): ], ) - result = itir_type_inference.infer(testee, offset_provider_type=mesh.offset_provider_type) + result = itir_type_inference.infer(testee, table_types=mesh.table_types) program_type = it_ts.ProgramType( params={"inp": float_edge_k_field, "out": float_vertex_k_field} @@ -493,7 +487,7 @@ def test_function_definition(): ], ) - result = itir_type_inference.infer(testee, offset_provider_type={}) + result = itir_type_inference.infer(testee, table_types={}) program_type = it_ts.ProgramType(params={"inp": float_i_field, "out": float_i_field}) assert result.type == program_type @@ -525,7 +519,7 @@ def test_fencil_with_nb_field_input(): ], ) - result = itir_type_inference.infer(testee, offset_provider_type=mesh.offset_provider_type) + result = itir_type_inference.infer(testee, table_types=mesh.table_types) stencil = result.body[0].expr.fun.args[0] assert stencil.expr.args[0].type == float64_list_type assert stencil.type.returns == float64_type @@ -550,7 +544,7 @@ def test_program_tuple_setat_short_target(): ], ) - result = itir_type_inference.infer(testee, offset_provider_type={}) + result = itir_type_inference.infer(testee, table_types={}) assert ( isinstance(result.body[0].expr.type, ts.TupleType) @@ -581,7 +575,7 @@ def test_program_setat_without_domain(): ], ) - result = itir_type_inference.infer(testee, offset_provider_type={}) + result = itir_type_inference.infer(testee, table_types={}) assert result.body[0].expr.type, ts.FieldType(dims=[IDim], dtype=float64_type) @@ -603,9 +597,7 @@ def test_if_stmt(): false_branch=[], ) - result = itir_type_inference.infer( - testee, offset_provider_type={}, allow_undeclared_symbols=True - ) + result = itir_type_inference.infer(testee, table_types={}, allow_undeclared_symbols=True) assert result.cond.type == bool_type assert result.true_branch[0].expr.type == float_i_field @@ -615,7 +607,7 @@ def test_as_fieldop_without_domain_nb_field_input(): testee = im.as_fieldop(stencil)(im.ref("inp1", float_vertex_v2e_field)) result = itir_type_inference.infer( - testee, offset_provider_type={V2EDim.tag: V2E}, allow_undeclared_symbols=True + testee, table_types={V2EDim.tag: V2E}, allow_undeclared_symbols=True ) assert result.type == ts.FieldType(dims=[Vertex], dtype=float64_list_type) assert result.fun.args[0].type.pos_only_args[0] == it_ts.IteratorType( @@ -632,7 +624,7 @@ def test_as_fieldop_without_domain_nb_field_input(): def test_comparison_with_non_scalar_rhs(rhs_type): testee = im.less(im.ref("a", int_type), im.ref("b", rhs_type)) with pytest.raises(AssertionError): - itir_type_inference.infer(testee, offset_provider_type={}, allow_undeclared_symbols=True) + itir_type_inference.infer(testee, table_types={}, allow_undeclared_symbols=True) def test_reinference(): diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py index c476dc4ee8..63e387a164 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_concat_where_expand_tuple_args.py @@ -44,9 +44,7 @@ def test_trivial(uids: utils.IDGeneratorPool): expected = im.make_tuple(im.concat_where(cond, "a", "b"), im.concat_where(cond, "c", "d")) - actual = concat_where.expand_tuple_args( - testee, offset_provider_type={}, allow_undeclared_symbols=True - ) + actual = concat_where.expand_tuple_args(testee, table_types={}, allow_undeclared_symbols=True) actual = collapse_tuple.CollapseTuple.apply( actual, allow_undeclared_symbols=True, within_stencil=False, uids=uids @@ -84,9 +82,7 @@ def test_nested(uids: utils.IDGeneratorPool): ), ) - actual = concat_where.expand_tuple_args( - testee, offset_provider_type={}, allow_undeclared_symbols=True - ) + actual = concat_where.expand_tuple_args(testee, table_types={}, allow_undeclared_symbols=True) actual = collapse_tuple.CollapseTuple.apply( actual, allow_undeclared_symbols=True, within_stencil=False, uids=uids diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py index a0357b9976..21d27f7ce8 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_cse.py @@ -24,7 +24,7 @@ class I(common.CartesianAxisIndex, kind=common.DimensionKind.HORIZONTAL): ... @pytest.fixture -def offset_provider_type(request): +def table_types(request): return {"I": I} @@ -154,7 +154,7 @@ def common_expr(): assert actual == expected -def test_if_can_deref_no_extraction(offset_provider_type, uids: utils.IDGeneratorPool): +def test_if_can_deref_no_extraction(table_types, uids: utils.IDGeneratorPool): # Test that a subexpression only occurring in one branch of an `if_` is not moved outside the # if statement. A case using `can_deref` is used here as it is common. @@ -175,12 +175,12 @@ def test_if_can_deref_no_extraction(offset_provider_type, uids: utils.IDGenerato ) actual = CommonSubexpressionElimination.apply( - testee, offset_provider_type=offset_provider_type, within_stencil=True, uids=uids + testee, table_types=table_types, within_stencil=True, uids=uids ) assert actual == expected -def test_if_can_deref_eligible_extraction(offset_provider_type, uids: utils.IDGeneratorPool): +def test_if_can_deref_eligible_extraction(table_types, uids: utils.IDGeneratorPool): # Test that a subexpression only occurring in both branches of an `if_` is moved outside the # if statement. A case using `can_deref` is used here as it is common. @@ -198,12 +198,12 @@ def test_if_can_deref_eligible_extraction(offset_provider_type, uids: utils.IDGe ) actual = CommonSubexpressionElimination.apply( - testee, offset_provider_type=offset_provider_type, within_stencil=True, uids=uids + testee, table_types=table_types, within_stencil=True, uids=uids ) assert actual == expected -def test_if_eligible_extraction(offset_provider_type, uids: utils.IDGeneratorPool): +def test_if_eligible_extraction(table_types, uids: utils.IDGeneratorPool): # Test that a subexpression only occurring in the condition of an `if_` is moved outside the # if statement. @@ -213,7 +213,7 @@ def test_if_eligible_extraction(offset_provider_type, uids: utils.IDGeneratorPoo expected = im.let("_cs_0", im.and_("a", "b"))(im.if_(im.and_("_cs_0", "_cs_0"), "c", "d")) actual = CommonSubexpressionElimination.apply( - testee, offset_provider_type=offset_provider_type, within_stencil=True, uids=uids + testee, table_types=table_types, within_stencil=True, uids=uids ) assert actual == expected diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py index bc1dbc3355..933915922a 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_dead_code_elimination.py @@ -53,7 +53,5 @@ def test_let_constant_foldable_if( input: itir.Expr, expected: itir.Expr, uids: utils.IDGeneratorPool ): input_program = program_factory(input) - inlined = dead_code_elimination.dead_code_elimination( - input_program, offset_provider_type={}, uids=uids - ) + inlined = dead_code_elimination.dead_code_elimination(input_program, table_types={}, uids=uids) assert inlined == program_factory(expected) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py index 8e9a0be03a..744bf8c817 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_expand_tuple_maps.py @@ -25,13 +25,13 @@ class IDim(common.CartesianAxisIndex): ... def _apply(expr: itir.Expr) -> itir.Expr: - return ExpandTupleMaps.apply(expr, uids=utils.IDGeneratorPool(), offset_provider_type={}) + return ExpandTupleMaps.apply(expr, uids=utils.IDGeneratorPool(), table_types={}) def _apply_and_collapse(expr: itir.Expr) -> itir.Expr: """Expand and then run the regular `CollapseTuple` pass, as happens in the pipeline.""" uids = utils.IDGeneratorPool() - result = ExpandTupleMaps.apply(expr, uids=uids, offset_provider_type={}) + result = ExpandTupleMaps.apply(expr, uids=uids, table_types={}) return CollapseTuple.apply( result, within_stencil=False, allow_undeclared_symbols=True, uids=uids ) @@ -84,7 +84,7 @@ def test_apply_creates_default_uids(_unary_fun): im.make_tuple(im.ref("a", i_field), im.ref("b", i_field)) ) - result = ExpandTupleMaps.apply(expr, uids=None, offset_provider_type={}) + result = ExpandTupleMaps.apply(expr, uids=None, table_types={}) assert expr.type is None assert result == im.let("_etm_0", im.make_tuple("a", "b"))( 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 5577e016be..0794ae3416 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 @@ -49,7 +49,7 @@ def test_trivial(uids: utils.IDGeneratorPool): d, )(im.ref("inp1", field_type), im.ref("inp2", field_type), im.ref("inp3", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, uids=uids ) assert actual == expected @@ -59,7 +59,7 @@ def test_trivial_literal(uids: utils.IDGeneratorPool): testee = im.op_as_fieldop("plus", d)(im.op_as_fieldop("multiplies", d)(1, 2), 3) expected = im.as_fieldop(im.lambda_()(im.plus(im.multiplies_(1, 2), 3)), d)() actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, uids=uids ) assert actual == expected @@ -78,7 +78,7 @@ def test_trivial_same_arg_twice(uids: utils.IDGeneratorPool): d, )(im.ref("inp1", field_type), im.ref("inp2", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -103,7 +103,7 @@ def test_tuple_arg(uids: utils.IDGeneratorPool): d, )() actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -123,7 +123,7 @@ def test_symref_used_twice(uids: utils.IDGeneratorPool): d, )("inp1", "inp2") actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -139,7 +139,7 @@ def test_no_inline(uids: utils.IDGeneratorPool): )(im.op_as_fieldop("plus", d2)(im.ref("inp1", field_type), im.ref("inp2", field_type))) actual = fuse_as_fieldop.FuseAsFieldOp.apply( testee, - offset_provider_type={}, + table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids, @@ -166,7 +166,7 @@ def test_staged_inlining(uids: utils.IDGeneratorPool): d, )(im.ref("a", field_type), im.ref("b", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -182,7 +182,7 @@ def test_make_tuple_fusion_trivial(uids: utils.IDGeneratorPool): d, )(im.ref("a", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) # simplify to remove unnecessary make_tuple call `{v[0], v[1]}(actual)` actual_simplified = ct.CollapseTuple.apply( @@ -202,7 +202,7 @@ def test_make_tuple_fusion_symref(uids: utils.IDGeneratorPool): d, )(im.ref("a", field_type), im.ref("b", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) # simplify to remove unnecessary make_tuple call actual_simplified = ct.CollapseTuple.apply( @@ -222,7 +222,7 @@ def test_make_tuple_fusion_symref_same_ref(uids: utils.IDGeneratorPool): d, )(im.ref("a", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) # simplify to remove unnecessary make_tuple call actual_simplified = ct.CollapseTuple.apply( @@ -247,7 +247,7 @@ def test_make_tuple_nested(uids: utils.IDGeneratorPool): d, )(im.ref("a", field_type), im.ref("b", field_type), im.ref("c", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) # simplify to remove unnecessary make_tuple call actual_simplified = ct.CollapseTuple.apply( @@ -289,7 +289,7 @@ def test_make_tuple_fusion_different_domains(uids: utils.IDGeneratorPool): ) ) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -326,7 +326,7 @@ def test_partial_inline(uids: utils.IDGeneratorPool): ) actual = fuse_as_fieldop.FuseAsFieldOp.apply( testee, - offset_provider_type={}, + table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids, @@ -353,7 +353,7 @@ def test_chained_fusion(uids: utils.IDGeneratorPool): d, )(im.ref("inp1", field_type), im.ref("inp2", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -374,7 +374,7 @@ def test_inline_as_fieldop_with_list_dtype(uids: utils.IDGeneratorPool): im.lambda_("inp")(im.call(im.call("reduce")("plus", 0))(im.deref("inp"))), d )(im.ref("inp", list_field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -385,7 +385,7 @@ def test_inline_into_scan(uids: utils.IDGeneratorPool): testee = im.as_fieldop(scan, d)(im.as_fieldop("deref")(im.ref("a", field_type))) expected = im.as_fieldop(scan, d)(im.ref("a", field_type)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected @@ -398,7 +398,7 @@ def test_no_inline_into_scan(uids: utils.IDGeneratorPool): scan = im.as_fieldop(scan_stencil, d)(im.ref("a", field_type)) testee = im.as_fieldop(im.lambda_("arg")(im.deref("arg")), d)(scan) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == testee @@ -411,6 +411,6 @@ def test_opage_arg_deduplication(uids: utils.IDGeneratorPool): d, )(im.index(IDim)) actual = fuse_as_fieldop.FuseAsFieldOp.apply( - testee, offset_provider_type={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids + testee, table_types={}, allow_undeclared_symbols=True, enable_cse=False, uids=uids ) assert actual == expected diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py index 0c39338b8d..945c5df1a0 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_global_tmps.py @@ -79,7 +79,7 @@ def test_trivial(uids: utils.IDGeneratorPool): ) ], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) testee = infer_domain.infer_program(testee, offset_provider=offset_provider) expected = program_factory( @@ -114,7 +114,7 @@ def test_trivial_let(uids: utils.IDGeneratorPool): ) ], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) testee = infer_domain.infer_program(testee, offset_provider=offset_provider) expected = program_factory( @@ -210,7 +210,7 @@ def test_top_level_if(uids: utils.IDGeneratorPool): ) ], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) testee = infer_domain.infer_program(testee, offset_provider=offset_provider) expected = program_factory( @@ -268,7 +268,7 @@ def test_nested_if(uids: utils.IDGeneratorPool): ) ], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) testee = infer_domain.infer_program(testee, offset_provider=offset_provider) expected = program_factory( @@ -336,7 +336,7 @@ def add_shifted(domain: itir.FunCall | None = None): ) ], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) testee = infer_domain.infer_program(testee, offset_provider=offset_provider) expected = program_factory( @@ -443,7 +443,7 @@ def add_shifted(domain: itir.FunCall | None = None): ) ], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) testee = infer_domain.infer_program(testee, offset_provider=offset_provider) expected = program_factory( @@ -530,7 +530,7 @@ def test_domain_preservation(uids: utils.IDGeneratorPool): ) ], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) expected = program_factory( params=[im.sym("inp", i_field_type), im.sym("out", i_field_type)], @@ -576,7 +576,7 @@ def test_non_scan_projector(uids: utils.IDGeneratorPool): ], body=[stmt], ) - testee = type_inference.infer(testee, offset_provider_type=offset_provider) + testee = type_inference.infer(testee, table_types=offset_provider) # make sure the statement actually has a projector projector, expr = ir_utils_misc.extract_projector(stmt.expr) diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py index f391363e76..341a73d88a 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_inline_scalar.py @@ -42,12 +42,12 @@ def program_factory(expr: itir.Expr) -> itir.Program: def test_simple(): testee = program_factory(im.let("a", 1)(im.op_as_fieldop("plus")("inp", "a"))) expected = program_factory(im.op_as_fieldop("plus")("inp", 1)) - actual = inline_scalar.InlineScalar.apply(testee, offset_provider_type={}) + actual = inline_scalar.InlineScalar.apply(testee, table_types={}) assert actual == expected def test_fo_inline_only(): scalar_expr = im.let("a", 1)(im.plus("a", "a")) testee = program_factory(im.as_fieldop(im.lambda_()(scalar_expr))()) - actual = inline_scalar.InlineScalar.apply(testee, offset_provider_type={}) + actual = inline_scalar.InlineScalar.apply(testee, table_types={}) assert actual == testee diff --git a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py index a745e08847..5aaf4705f5 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_prune_casts.py @@ -20,7 +20,7 @@ def test_prune_casts_simple(): x_ref = im.ref("x", ts.ScalarType(kind=ts.ScalarKind.FLOAT32)) y_ref = im.ref("y", ts.ScalarType(kind=ts.ScalarKind.FLOAT64)) testee = im.plus(im.cast_(x_ref, "float64"), im.cast_(y_ref, "float64")) - testee = type_inference.infer(testee, offset_provider_type={}, allow_undeclared_symbols=True) + testee = type_inference.infer(testee, table_types={}, allow_undeclared_symbols=True) expected = im.plus(im.cast_(x_ref, "float64"), y_ref) actual = PruneCasts.apply(testee) @@ -34,7 +34,7 @@ def test_prune_casts_fieldop(): im.cast_as_fieldop("float64")(x_ref), im.cast_as_fieldop("float64")(y_ref), ) - testee = type_inference.infer(testee, offset_provider_type={}, allow_undeclared_symbols=True) + testee = type_inference.infer(testee, table_types={}, allow_undeclared_symbols=True) expected = im.op_as_fieldop("plus")( im.cast_as_fieldop("float64")(x_ref), 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 ae5c1186fe..367a5cf6c6 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 @@ -135,12 +135,10 @@ def _expected(red, max_neighbors, has_skip_values, shifted_arg=0): def test_basic(basic_reduction, has_skip_values, uids: utils.IDGeneratorPool): expected = _expected(basic_reduction, 3, has_skip_values) - offset_provider_type = { + table_types = { Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=has_skip_values) } - actual = UnrollReduce.apply( - basic_reduction, offset_provider_type=offset_provider_type, uids=uids - ) + actual = UnrollReduce.apply(basic_reduction, table_types=table_types, uids=uids) assert actual == expected @@ -149,11 +147,11 @@ def test_reduction_with_shift_on_second_arg( ): expected = _expected(reduction_with_shift_on_second_arg, 1, has_skip_values, 1) - offset_provider_type = { + table_types = { Dim.tag: dummy_connectivity_type(max_neighbors=1, has_skip_values=has_skip_values) } actual = UnrollReduce.apply( - reduction_with_shift_on_second_arg, offset_provider_type=offset_provider_type, uids=uids + reduction_with_shift_on_second_arg, table_types=table_types, uids=uids ) assert actual == expected @@ -161,10 +159,8 @@ def test_reduction_with_shift_on_second_arg( def test_reduction_with_if(reduction_if, uids: utils.IDGeneratorPool): expected = _expected(reduction_if, 2, False) - offset_provider_type = { - Dim.tag: dummy_connectivity_type(max_neighbors=2, has_skip_values=False) - } - actual = UnrollReduce.apply(reduction_if, offset_provider_type=offset_provider_type, uids=uids) + table_types = {Dim.tag: dummy_connectivity_type(max_neighbors=2, has_skip_values=False)} + actual = UnrollReduce.apply(reduction_if, table_types=table_types, uids=uids) assert actual == expected @@ -173,20 +169,20 @@ def test_reduction_with_irrelevant_full_shift( ): expected = _expected(reduction_with_irrelevant_full_shift, 3, False) - offset_provider_type = { + table_types = { Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=False), "IrrelevantDim": dummy_connectivity_type( max_neighbors=1, has_skip_values=True ), # different max_neighbors and skip value to trigger error } actual = UnrollReduce.apply( - reduction_with_irrelevant_full_shift, offset_provider_type=offset_provider_type, uids=uids + reduction_with_irrelevant_full_shift, table_types=table_types, uids=uids ) assert actual == expected @pytest.mark.parametrize( - "offset_provider_type", + "table_types", [ { Dim.tag: dummy_connectivity_type(max_neighbors=3, has_skip_values=False), @@ -202,10 +198,6 @@ def test_reduction_with_irrelevant_full_shift( }, ], ) -def test_reduction_with_incompatible_shifts( - reduction_with_incompatible_shifts, offset_provider_type, uids -): +def test_reduction_with_incompatible_shifts(reduction_with_incompatible_shifts, table_types, uids): with pytest.raises(RuntimeError, match="incompatible"): - UnrollReduce.apply( - reduction_with_incompatible_shifts, offset_provider_type=offset_provider_type, uids=uids - ) + UnrollReduce.apply(reduction_with_incompatible_shifts, table_types=table_types, uids=uids) diff --git a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_itir_to_gtfn_ir.py b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_itir_to_gtfn_ir.py index 7da31a35bf..fad61efd8e 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_itir_to_gtfn_ir.py +++ b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_itir_to_gtfn_ir.py @@ -21,7 +21,7 @@ def test_funcall_to_op(): ) actual = it2gtfn.GTFN_lowering( - grid_type=gtx.GridType.CARTESIAN, offset_provider_type={}, column_axis=None + grid_type=gtx.GridType.CARTESIAN, table_types={}, column_axis=None ).visit(testee) assert expected == actual @@ -32,7 +32,7 @@ def test_unapplied_funcall_to_function_object(): expected = gtfn_ir.SymRef(id="plus") actual = it2gtfn.GTFN_lowering( - grid_type=gtx.GridType.CARTESIAN, offset_provider_type={}, column_axis=None + grid_type=gtx.GridType.CARTESIAN, table_types={}, column_axis=None ).visit(testee) assert expected == actual diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py index faea3c7764..1a02371357 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace.py @@ -208,7 +208,7 @@ def verify_testee(): def test_dace_fastcall_with_connectivity(unstructured_case, monkeypatch): """Test reuse of SDFG arguments between program calls by means of SDFG fastcall API.""" - connectivity_E2V = unstructured_case.offset_provider[E2VDim.tag].asnumpy() + connectivity_E2V = unstructured_case.offset_provider[E2V].asnumpy() @gtx.field_operator def testee(a: cases.VField) -> cases.EField: diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index ff1ec0addd..b3cf628a8b 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -379,7 +379,7 @@ def testee(a: cases.VField, b: cases.VField): ), ) - SIMPLE_MESH = cases_utils.simple_mesh(None) + SIMPLE_MESH = cases_utils.ir_level(cases_utils.simple_mesh(None)) offset_provider = SIMPLE_MESH.offset_provider test_case = cases.Case.from_mesh_descriptor(SIMPLE_MESH, backend=backend, allocator=backend) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py index 1bec200fad..3a654e76be 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_translation.py @@ -29,6 +29,7 @@ ) from gt4py.next.type_system import type_specifications as ts +from next_tests.integration_tests import cases_utils from next_tests.integration_tests.cases_utils import ( V2E, Edge, @@ -80,7 +81,7 @@ def _translate_gtir_to_sdfg( @pytest.mark.parametrize("has_unit_stride", [False, True]) @pytest.mark.parametrize("disable_field_origin", [False, True]) def test_find_constant_symbols(has_unit_stride, disable_field_origin): - SKIP_VALUE_MESH = skip_value_mesh(None) + SKIP_VALUE_MESH = cases_utils.ir_level(skip_value_mesh(None)) ir = itir.Program( id="find_constant_symbols_sdfg", @@ -94,7 +95,7 @@ def test_find_constant_symbols(has_unit_stride, disable_field_origin): itir.SetAt( expr=im.as_fieldop( im.lambda_("it")(im.reduce("plus", im.literal_from_value(1.0))(im.deref("it"))) - )(im.as_fieldop_neighbors(V2E.value, "x")), + )(im.as_fieldop_neighbors(V2E.Local.tag, "x")), domain=im.get_field_domain(gtx_common.GridType.UNSTRUCTURED, "y", VFTYPE.dims), target=itir.SymRef(id="y"), ) @@ -112,7 +113,7 @@ def test_find_constant_symbols(has_unit_stride, disable_field_origin): constant_symbols = dace_wf_translation.find_constant_symbols( ir=ir, sdfg=sdfg, - offset_provider_type=SKIP_VALUE_MESH.offset_provider_type, + table_types=SKIP_VALUE_MESH.table_types, disable_field_origin_on_program_arguments=disable_field_origin, unstructured_horizontal_has_unit_stride=has_unit_stride, ) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py index a1fb215882..9d298408aa 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_gtir_to_sdfg.py @@ -26,6 +26,7 @@ from gt4py.next.iterator.transforms import pass_manager from gt4py.next.type_system import type_specifications as ts +from next_tests.integration_tests import cases_utils from next_tests.integration_tests.cases_utils import ( E2VDim, C2VDim, @@ -72,8 +73,8 @@ def allow_view_arguments(): IOff = im.cartesian_offset(IDim, IDim) # Cartesian shifts are self-describing (`CartesianOffset`), so no offset provider entry is needed. CARTESIAN_OFFSETS: dict = {} -SIMPLE_MESH: MeshDescriptor = simple_mesh(None) -SKIP_VALUE_MESH: MeshDescriptor = skip_value_mesh(None) +SIMPLE_MESH: MeshDescriptor = cases_utils.ir_level(simple_mesh(None)) +SKIP_VALUE_MESH: MeshDescriptor = cases_utils.ir_level(skip_value_mesh(None)) SIZE_TYPE = ts.ScalarType(ts.ScalarKind.INT32) FSYMBOLS = dict( **{gtx_dace_args.range_start_symbol("w", IDim).name: 0}, @@ -194,8 +195,8 @@ def build_dace_sdfg( offset_provider=offset_provider, symbolic_domain_sizes=pass_manager._max_domain_range_sizes(offset_provider), ) - offset_provider_type = gtx_common.offset_provider_to_type(offset_provider) - return dace_lowering.lower_program_to_sdfg(ir, offset_provider_type, column_axis=KDim) + table_types = gtx_common.offset_provider_to_type(offset_provider) + return dace_lowering.lower_program_to_sdfg(ir, table_types, column_axis=KDim) def apply_margin_on_field_domain( diff --git a/tests/next_tests/unit_tests/test_common.py b/tests/next_tests/unit_tests/test_common.py index f307f9d5b3..b9e993cf07 100644 --- a/tests/next_tests/unit_tests/test_common.py +++ b/tests/next_tests/unit_tests/test_common.py @@ -24,6 +24,7 @@ DimensionIndex, LocalDimensionIndex, DimensionKind, + LocalDimensionIndex, Domain, Infinity, UnitRange, diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 8e47de451f..a62c56ecc3 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -47,6 +47,10 @@ class Local(LocalDimensionIndex): ... class LsqCoeff(LocalDimensionIndex, size=3): ... +class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local = V2E.Local + + def _declare(source: str) -> dict: """ Run `source` as the body of a throwaway module. @@ -474,10 +478,6 @@ def test_shift_type_str(self): ) assert str(ts.ShiftType(codomain=KDim, domain=(KDim,))) == f"Shift[{KDim} -> {KDim}]" - def test_field_offset_is_derived_once(self): - assert V2E.__gt_field_offset__() is V2E.__gt_field_offset__() - assert V2E.__gt_field_offset__().value == V2E.Local.tag - def test_neighbor_index_accepts_numpy_integers(self): from gt4py.next import constructors, embedded @@ -486,14 +486,9 @@ def test_neighbor_index_accepts_numpy_integers(self): codomain=Edge, data=np.array([[0, 1, 2, 3], [1, 2, 3, 0]]), ) - with embedded.context.update(offset_provider={V2E.Local.tag: table}): + with embedded.context.update(offset_provider={V2E: 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 @@ -543,6 +538,18 @@ def test_grid_type_deduction(self): transform_utils._deduce_grid_type(common.GridType.CARTESIAN, [V2E]) +def _table(domain=(Vertex, V2E.Local), codomain=Edge, data=((0, 1, 2, 3), (1, 2, 3, 0))): + from gt4py.next import constructors + + data = np.array(data) + return constructors.as_connectivity( + domain=dict(zip(domain, data.shape)), + codomain=codomain, + data=data, + skip_value=common._DEFAULT_SKIP_VALUE, + ) + + def test_redefined_declaration_with_an_adopted_local(monkeypatch): """Re-running a cell must re-own the adopted local dimension, not become a sharer.""" import sys @@ -586,6 +593,84 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): common.local_dimension_of(NeighborConnectivity) +class TestOffsetProvider: + def test_class_keys_become_tags(self): + table = _table() + assert common.as_tag_keyed_offset_provider({V2E: table}) == {V2E.Local.tag: table} + + def test_tag_keys_pass_through(self): + provider = {V2E.Local.tag: _table()} + assert common.as_tag_keyed_offset_provider(provider) is provider + + def test_bare_name_is_rejected(self): + with pytest.raises(TypeError, match="keyed by 'NeighborConnectivity' declarations"): + common.as_tag_keyed_offset_provider({"V2E": _table()}) + + def test_non_string_key_is_rejected(self): + with pytest.raises(TypeError, match="keyed by 'NeighborConnectivity' declarations"): + common.as_tag_keyed_offset_provider({V2E: _table(), Vertex: _table()}) + + def test_binding_twice_is_rejected(self): + with pytest.raises(ValueError, match="twice"): + common.as_tag_keyed_offset_provider({V2E: _table(), V2E.Local.tag: _table()}) + + def test_check_accepts_matching_tables(self): + common.check_offset_provider({V2E.Local.tag: _table()}) + table = _table() + # the types alone, as an ahead-of-time compilation gives them + common.check_offset_provider( + { + V2E: NeighborTableType( + connectivity=V2E, + dtype=table.dtype, + skip_value=table.skip_value, + max_neighbors=4, + ) + } + ) + + def test_check_rejects_mismatching_tables(self): + with pytest.raises(ValueError, match="does not match its declaration"): + common.check_offset_provider({V2E.Local.tag: _table(codomain=Vertex)}) + + def test_sharing_connectivity_is_keyed_and_checked_by_its_own_tag(self): + provider = common.as_tag_keyed_offset_provider({V2E: _table(), V2EShared: _table()}) + assert set(provider) == {V2E.Local.tag, V2EShared.tag} + common.check_offset_provider(provider) + with pytest.raises(ValueError, match="'V2EShared' does not match its declaration"): + common.check_offset_provider({V2EShared.tag: _table(codomain=Vertex)}) + + def test_sharing_connectivities_need_the_same_skip_positions(self): + owner = _table(data=((0, 1, 2, 3), (1, 2, 3, 0))) + consistent = _table(data=((3, 2, 1, 0), (0, 3, 2, 1))) + inconsistent = _table(data=((3, 2, 1, -1), (0, 3, 2, 1))) + common.check_offset_provider({V2E: owner, V2EShared: consistent}, deep=True) + with pytest.raises(ValueError, match="different neighbor structure"): + common.check_offset_provider({V2E: owner, V2EShared: inconsistent}, deep=True) + # the call path does not read the tables + common.check_offset_provider({V2E: owner, V2EShared: inconsistent}) + + def test_the_connectivity_s_own_tag_is_rejected(self): + with pytest.raises(ValueError, match="whose key is the declaration itself"): + common.check_offset_provider({V2E.tag: _table()}) + + def test_a_checked_provider_is_remembered(self): + provider = {V2E: _table(codomain=Vertex)} + with pytest.raises(ValueError, match="does not match its declaration"): + common.check_offset_provider(provider) + # and a provider that passes is not checked twice + good = {V2E: _table()} + common.check_offset_provider(good) + common.check_offset_provider(good) + + def test_check_skips_undeclared_tags(self): + common.check_offset_provider({"some.hand.written.tag": _table()}) + + def test_missing_connectivity_error_is_actionable(self): + with pytest.raises(KeyError, match="keyed by 'NeighborConnectivity' declarations"): + common.get_offset({}, V2E.Local.tag) + + class TestConnectivityKeyOver: def _type(self, connectivity): return _table_type(domain=(connectivity.domain, common.local_dimension_of(connectivity))) @@ -626,6 +711,39 @@ def test_nothing_bound(self): common.connectivity_key_over({E2V.offset_tag: self._type(E2V)}, V2E.Local) +def test_redefined_declaration_resolves_to_the_new_class(monkeypatch): + """Re-running a notebook cell redefines declarations under the same names.""" + import sys + import types as pytypes + + module = pytypes.ModuleType("_redefined_connectivity_module") + monkeypatch.setitem(sys.modules, module.__name__, module) + source = textwrap.dedent( + """ + from gt4py.next.common import DimensionIndex, LocalDimensionIndex, NeighborConnectivity + + class V(DimensionIndex): ... + class E(DimensionIndex): ... + class V2E(NeighborConnectivity[V, E], max_neighbors={n}): + class Local(LocalDimensionIndex): ... + """ + ) + exec(source.format(n=4), module.__dict__) + old = module.V2E + common.check_offset_provider({old: _table(domain=(module.V, old.Local), codomain=module.E)}) + + exec(source.format(n=2), module.__dict__) + new = module.V2E + assert common.resolve(new.Local.tag) is new.Local + common.check_offset_provider( + {new: _table(domain=(module.V, new.Local), codomain=module.E, data=((0, 1), (1, 0)))} + ) + with pytest.raises(ValueError, match="was the declaration redefined"): + common.check_neighbor_table( + new, _table(domain=(old.domain, old.Local), codomain=module.E, data=((0, 1), (1, 0))) + ) + + def test_the_const_list_dimension_cannot_be_adopted(): with pytest.raises(TypeError, match="cannot adopt"): _declare( 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 f006282e53..cdd9f53a3e 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 @@ -16,6 +16,7 @@ DimensionIndex, LocalDimensionIndex, DimensionKind, + LocalDimensionIndex, ) from gt4py.next.type_system import type_info, type_specifications as ts from gt4py.next.ffront import type_specifications as ts_ffront diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index a65d3e7d09..3519a3eb82 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -345,7 +345,7 @@ 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]" + main:17:15: note: "Field[Dims[Vertex, Local], float].__call__" has type "def __call__(self, index_field: Connectivity[Any, Any] | type[NeighborConnectivity[Any, Any]], *args: Connectivity[Any, Any] | type[NeighborConnectivity[Any, Any]]) -> Field[Any, Any]" - case: neighbor_connectivity_generic_local main: |