Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
158 changes: 158 additions & 0 deletions docs/development/ADRs/next/0029-Connectivities_As_Types.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
---
tags: []
---

# Connectivities as Types

- **Status**: proposed
- **Authors**: Enrique González Paredes (@egparedes)
- **Created**: 2026-09-21
- **Updated**: 2026-09-21

A neighbor connectivity is declared as a **class**, and its local dimension as a
class **nested** in it:

```python
class V2E(gtx.NeighborConnectivity[Vertex, Edge], max_neighbors=6, min_neighbors=5):
class Local(gtx.LocalDimensionIndex): ...


@gtx.field_operator
def f(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]:
return neighbor_sum(a(V2E), axis=V2E.Local) + a(V2E[0])
```

The declaration is written in DSL code, owns its local dimension, and states the
constraints a neighbor table bound to it has to satisfy. It holds no data. It builds on [ADR 0028](0028-Dimensions_As_Nominal_Types.md): the
connectivity, like a dimension, is identified by its type, and `V2E.Local` is an
ordinary dimension class with the tag `<module>.V2E.Local`.

## Context

An unstructured connectivity used to be spelled by four independently authored
names that had to agree, none of them checked against the others: the
`FieldOffset` tag, the Python variable it was bound to, the local dimension's
name and the offset-provider key. The `V2EDim = Dimension("V2E")` convention made
all four equal, which hid which one each execution path actually used; the
regression tests in `test_offset_dimensions_names.py` break the convention one
name at a time. Nothing tied a local dimension to the table it indexes, so the
backends recovered that link by string equality, and the table's shape, codomain
and skip values were never checked against the `FieldOffset` declaration.

## Decision

### The declaration

- `NeighborConnectivity[Origin, Codomain]` is a PEP 695 generic whose subclasses
are declarations: for each `Origin` element, a list of `Codomain` neighbors. Its
metaclass, `ConnectivityMeta`, forbids instantiation.
- The local dimension is the nested class `Local`, a subclass of
`LocalDimensionIndex`. Declaring it is required, and `NeighborConnectivity`
sets `Local.owner` to the connectivity when the class is created. A local
dimension can have at most one owner; a declaration redefined under the same
name (a re-run notebook cell) takes ownership over again, and for a local
dimension adopted rather than nested, the first declaration wins.
- A local dimension with no table, such as the coefficient axis of a fixed-size
stencil, is declared on its own: `class LsqCoeff(LocalDimensionIndex, size=3)`.
Its `owner` is `None`. A declaration can also *adopt* such a module-level local
dimension, written `Local: TypeAlias = LsqCoeff`, which then keeps its own tag.
- A connectivity can *share* another one's local dimension,
`Local: TypeAlias = C2E.Local`.
This is the flattened sparse pattern, e.g. cell-to-cell-edge (`C2CE: Cell -> CellEdge`) indexing the same neighbor axis as `C2E`, so that its results
combine with `C2E`-shaped sparse fields. The owner stays `C2E`, and the
neighbor counts and skip-value structure are the owner's.
- `max_neighbors` and `min_neighbors` are optional class keywords, not type
parameters: Python has no integer type parameters, and nothing static needs
the count. A declared count is a constraint on the bound table; an undeclared
one is taken from the table. `min_neighbors < max_neighbors` means that the
table must use skip values.
- `common.check_neighbor_table(V2E, table)` checks a table, or just its type
(which is all an ahead-of-time compilation has), against the declaration:
the domain is `(Origin, V2E.Local)`, the codomain is `Codomain`, the dtype is
integral, and the neighbor counts and skip values agree. Skip values are
checked on the table's type: a table with a `skip_value` counts as having skip
values whether or not an entry uses it. The check is explicit for as long as
offset providers are keyed by tag strings, since nothing then connects a
provider entry to a declaration; it becomes automatic with class-keyed
providers.

### `NeighborConnectivity` is not a `Connectivity`

`common.Connectivity` is a *data* protocol (`ndarray`, `domain`, `asnumpy`); a
declaration holds no data. The neighbor table stays a `Connectivity`
implementation, and the declaration is only the type the table is checked
against. `Field.premap` and `Field.__call__` accept either, as they already
accepted a `FieldOffset`, which is not a `Connectivity` either.

### `LocalDimensionIndex` subclasses `DimensionIndex`

A separate root would force every `type[DimensionIndex]` annotation in the tree
(`ts.FieldType.dims`, `Domain`, `ConnectivityType.domain`, ...) to widen, and
would then accept local dimensions wherever a primary one is meant anyway. The
tree already distinguishes local dimensions by a runtime `kind` check, so it
keeps doing so; generic constructors whose parameter must be a primary dimension
(`NeighborConnectivity[Origin, Codomain]`, `Staggered[D]`) check it at runtime.

### `Local` is not annotated anywhere

Neither `NeighborConnectivity` nor `ConnectivityMeta` annotates `Local`, and
that is load-bearing: an annotation makes a declaration's `Local` a *variable*
for the checkers, so `Field[Dims[Vertex, V2E.Local], float]` is rejected by
pyright ("Variable not allowed in type expression") for a nested `Local`, and by
mypy ("not valid as a type") for an adopted or shared one. A real nested `Local`
on the base is not an option either: pyright reports an incompatible override in
every declaration. With no annotation, all three spellings are types for both
checkers, which `typing_tests/pyright_probes.py` pins for pyright and
`typing_tests/test_next.yaml` for mypy.

The cost is that `conn.Local` is not an attribute the checkers know for a
*generic* `conn`. Library code reads it through `common.local_dimension_of(conn)`
instead, and code that has to name a local dimension generically uses a
`TypeVar` bound to `LocalDimensionIndex`. Writing an adopted or shared local as
`Local: TypeAlias = ...` (rather than a plain assignment) is what keeps mypy
treating it as a type.

### Frontend integration

A declaration is typed like the `FieldOffset` it replaces: `V2E.__gt_type__()`
is the `ts.OffsetType` of the derived offset `(Codomain -> (Origin, Local))`,
whose tag is the connectivity's `offset_tag`:

- **the local dimension's tag**, `V2E.Local.tag`, for the connectivity that
declares it. This is the single string that shifts, neighbor reductions and
sparse arguments already use to find the table in the offset provider, so
existing backends need no change.
- **its own tag**, `C2CE.tag`, for a connectivity that shares another one's local
dimension, since the local dimension's tag already names the owner's table.
Shifts find the table by that tag; reductions and sparse arguments still find
the owner's table by the local dimension's tag, for its neighbor structure, so
the owner has to be bound too.
`V2E.Local` inside DSL code types as that local dimension, and
`FieldOffset.Local` names the same thing on a legacy offset, so the spelling
works for both. The other frontend touch points treat the class like the
`FieldOffset` it derives: grid-type deduction (`transform_utils`, `past_to_itir`)
counts it as unstructured, and embedded `premap` accepts it. `V2E[i]` subscripts the
metaclass, which forwards type-parameter subscription (`NeighborConnectivity[V, E]`) to `__class_getitem__`, since a metaclass `__getitem__` shadows it.

## Consequences

- An unstructured connectivity is spelled once. The provider key, the offset tag
and the local dimension are all derived from the declaration.
- A table bound to a connectivity can be checked against its declaration.
- `V2E.Local` in DSL code is resolved from the offset type, because the type of
`V2E` is not the class.
- A declaration is fingerprinted by its name *and* its declared dimensions and
counts, so redefining it under the same name (e.g. re-running a notebook
cell) does not reuse artifacts compiled for the old declaration.
- `FieldOffset` remains during migration; a `FieldOffset` and a
`NeighborConnectivity` sharing a local dimension are interchangeable.

## Alternatives considered

- **The local dimension generated by the metaclass**, e.g. `V2E.Local` created
from `V2E`'s name. Type checkers cannot see a generated class, so it could not
be used in `Field[Dims[Vertex, V2E.Local], ...]`.
- **Neighbor counts as type parameters.** Python has no integer type parameters,
and a `Literal[6]` argument would add a type parameter nothing statically uses.
- **`NeighborConnectivity` as a `Connectivity` subclass.** Mixes the
declaration with the data protocol; see above.
1 change: 1 addition & 0 deletions docs/development/ADRs/next/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ Writing a new ADR is simple:
- [0023 - Fingerprinting](0023-Fingerprinting.md)
- [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md)
- [0028 - Dimensions as Nominal Types](0028-Dimensions_As_Nominal_Types.md)
- [0029 - Connectivities as Types](0029-Connectivities_As_Types.md)

### Frontend and Parsing #frontend

Expand Down
44 changes: 22 additions & 22 deletions docs/user/next/QuickstartGuide.md
Original file line number Diff line number Diff line change
Expand Up @@ -218,28 +218,28 @@ edge_values = gtx.as_field([EdgeDim], np.zeros((12,)))

+++

You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _field offset_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements.
You can transform fields (or tuples of fields) over one domain to another domain by using the call operator of the source field with a _connectivity_ as argument. This transform uses the connectivity between the source and target domains to find the values of adjacent mesh elements.

To understand this transform, you can look at the edge-to-cell connectivity table `edge_to_cell_table` listed above. This table has the same shape as the output of the transform, that is, one dimension over the edges and another _local_ dimension. The table stores indices into a field over cells, the transform essentially gives you another field where the indices have been replaced with the values in the cell field at the corresponding indices.

Another way to look at it is that transform uses the edge-to-cell connectivity to look up all the cell neighbors of edges, and associates the values of those neighbor cells with each edge.

You can use the field offset `E2C` below to transform a field over cells to a field over edges using the edge-to-cell connectivities:
You can use the connectivity `E2C` declared below to transform a field over cells to a field over edges using the edge-to-cell connectivities. It is declared as a class: for each edge (`EdgeDim`), a list of neighbor cells (`CellDim`), indexed by its nested local dimension `E2C.Local`:

```{code-cell} ipython3
class E2CDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...
E2C = gtx.FieldOffset(E2CDim.tag, source=CellDim, target=(EdgeDim, E2CDim))
class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]):
class Local(gtx.LocalDimensionIndex): ...
```

The field offset is named by its local dimension's `tag`, and the offset provider below is keyed by the same `tag`, so all three refer to one connectivity. Note that the field offset does not contain the actual connectivity table, that's provided through an _offset provider_:
Note that the declaration does not contain the actual connectivity table, that's provided through an _offset provider_, keyed by the local dimension's `tag`:

```{code-cell} ipython3
E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2CDim], codomain=CellDim, data=edge_to_cell_table, skip_value=-1)
E2C_offset_provider = gtx.as_connectivity([EdgeDim, E2C.Local], codomain=CellDim, data=edge_to_cell_table, skip_value=-1)
```

The field operator `nearest_cell_to_edge` below shows an example of applying this transform. There is a little twist though: the subscript in `E2C[0]` means that only the value of the first connected cell is taken, the second (if exists) is ignored.

Pay attention to the syntax where the field offset `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program.
Pay attention to the syntax where the connectivity `E2C` can be freely accessed in the field operator, but the offset provider `E2C_offset_provider` is passed in a dictionary to the program.

```{code-cell} ipython3
@gtx.field_operator
Expand All @@ -250,7 +250,7 @@ def nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64]) -> gtx.
def run_nearest_cell_to_edge(cell_values: gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]):
nearest_cell_to_edge(cell_values, out=out)

run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider})
run_nearest_cell_to_edge(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider})

print("0th adjacent cell's value: {}".format(edge_values.asnumpy()))
```
Expand All @@ -265,19 +265,19 @@ Running the above snippet results in the following edge field:

#### Using reductions on connected mesh elements

Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2CDim) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2CDim` dimension.
Similarly to the previous example, the output is once again a field on edges. The difference is that this field operator does not take the first column of the transformed field, but sums the columns. In other words, the result is the sum of all the cells adjacent to an edge. You can achieve this by first transforming the cell field to a field over the cell neighbors of edges (i.e. a field of dimensions Edge × E2C.Local) using `cells(E2C)`, then calling the `neighbor_sum` builtin function to sum along the `E2C.Local` dimension.

```{code-cell} ipython3
@gtx.field_operator
def sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64]) -> gtx.Field[Dims[EdgeDim], float64]:
# type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2CDim], float64]
return neighbor_sum(cells(E2C), axis=E2CDim)
# type of cells(E2C) is gtx.Field[Dims[EdgeDim, E2C.Local], float64]
return neighbor_sum(cells(E2C), axis=E2C.Local)

@gtx.program
def run_sum_adjacent_cells(cells : gtx.Field[Dims[CellDim], float64], out : gtx.Field[Dims[EdgeDim], float64]):
sum_adjacent_cells(cells, out=out)

run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2CDim.tag: E2C_offset_provider})
run_sum_adjacent_cells(cell_values, edge_values, offset_provider={E2C.Local.tag: E2C_offset_provider})

print("sum of adjacent cells: {}".format(edge_values.asnumpy()))
```
Expand Down Expand Up @@ -376,13 +376,13 @@ print("where nested tuple return: {}".format(((result_1.asnumpy(), result_2.asnu

#### Implementing the pseudo-laplacian

As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to define the local dimension, the field offset, and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before:
As explained in the section outline, the pseudo-laplacian needs the cell-to-edge connectivities as well in addition to the edge-to-cell connectivities. Though the connectivity table has been filled in above, you still need to declare the connectivity and the offset provider that describe how to use the connectivity table. The procedure is identical to the edge-to-cell connectivity from before:

```{code-cell} ipython3
class C2EDim(gtx.DimensionIndex, kind=gtx.DimensionKind.LOCAL): ...
C2E = gtx.FieldOffset(C2EDim.tag, source=EdgeDim, target=(CellDim, C2EDim))
class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]):
class Local(gtx.LocalDimensionIndex): ...

C2E_offset_provider = gtx.as_connectivity([CellDim, C2EDim], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1)
C2E_offset_provider = gtx.as_connectivity([CellDim, C2E.Local], codomain=EdgeDim, data=cell_to_edge_table, skip_value=-1)
```

**Weights of edge differences:**
Expand Down Expand Up @@ -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.
Expand All @@ -422,17 +422,17 @@ 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:

```{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)

Expand All @@ -441,7 +441,7 @@ result_pseudo_lap = gtx.as_field([CellDim], np.zeros(shape=(6,)))
run_pseudo_laplacian(cell_values,
edge_weight_field,
result_pseudo_lap,
offset_provider={E2CDim.tag: E2C_offset_provider, C2EDim.tag: C2E_offset_provider})
offset_provider={E2C.Local.tag: E2C_offset_provider, C2E.Local.tag: C2E_offset_provider})

print("pseudo-laplacian: {}".format(result_pseudo_lap.asnumpy()))
```
Expand All @@ -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)
```

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,7 @@
" A,\n",
" edge_orientation,\n",
" out=divergence_gt4py,\n",
" offset_provider={C2E.value: c2e_connectivity},\n",
" offset_provider={C2E.Local.tag: c2e_connectivity},\n",
" )\n",
"\n",
" assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@
" A,\n",
" edge_orientation,\n",
" out=divergence_gt4py,\n",
" offset_provider={C2E.value: c2e_connectivity},\n",
" offset_provider={C2E.Local.tag: c2e_connectivity},\n",
" )\n",
"\n",
" assert np.allclose(divergence_gt4py.asnumpy(), divergence_ref)"
Expand Down
Loading
Loading