diff --git a/docs/development/ADRs/next/0030-Connectivities_As_Types.md b/docs/development/ADRs/next/0030-Connectivities_As_Types.md index a5c88f32b5..fae2309732 100644 --- a/docs/development/ADRs/next/0030-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0030-Connectivities_As_Types.md @@ -248,6 +248,17 @@ 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. +### Positions in a neighborhood + +`MultiDimensionIndex[D, *Ls]` is a position in the product of a primary +dimension and local dimensions, e.g. `MultiDimensionIndex(Vertex(3), V2E.Local(1))`, the second neighbor of vertex 3. It is a tuple of indices, so it +indexes a field or a table directly; since a `TypeVarTuple` cannot carry a bound, +the constructor checks the shape. It is a typed convenience for users; nothing in the toolchain requires it. +The iterator-level embedded execution keys its +positions by dimension classes instead of tag strings, and steps along a field's +local dimension with an explicit `SparseAxis(dim)`. Named offsets still +arrive as IR strings and find their tables through the tag-keyed provider. + ## Consequences - An unstructured connectivity is spelled once. The provider key, the offset tag diff --git a/src/gt4py/next/__init__.py b/src/gt4py/next/__init__.py index 00803ce442..5969adb2bc 100644 --- a/src/gt4py/next/__init__.py +++ b/src/gt4py/next/__init__.py @@ -34,6 +34,7 @@ Field, GridType, LocalDimensionIndex, + MultiDimensionIndex, NeighborConnectivity, Staggered, UnitRange, @@ -127,6 +128,7 @@ "CartesianAxisIndex", "DimensionKind", "LocalDimensionIndex", + "MultiDimensionIndex", "NeighborConnectivity", "Staggered", "resolve", diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index e3159b8b68..874041964a 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -2691,3 +2691,67 @@ def _check_shared_local_dimensions( " local dimension must have the same number of neighbors, and skip values at" " the same positions." ) + + +class MultiDimensionIndex[D: DimensionIndex, *Ls](tuple[D, *Ls]): + """ + A position in the product of a primary dimension and local dimensions. + + For example the entry `(Vertex(3), V2E.Local(1))` of the table of `V2E`: the second neighbor + of vertex 3. It is a tuple of indices, so it indexes a field or a neighbor table directly, and + it compares and hashes like the plain tuple; tuple operations such as slicing return plain + tuples. A user-facing, typed index: nothing in the toolchain requires it. + + At least one local index is required, and a local dimension owned by a connectivity must be + one of the neighbors of the primary index's dimension. + + Examples: + >>> class Vertex(DimensionIndex): ... + >>> class Edge(DimensionIndex): ... + >>> class V2E(NeighborConnectivity[Vertex, Edge]): + ... class Local(LocalDimensionIndex): ... + >>> position = MultiDimensionIndex(Vertex(3), V2E.Local(1)) + >>> position + MultiDimensionIndex(Vertex=3, V2E.Local=1) + >>> position.dims == (Vertex, V2E.Local) + True + """ + + __slots__ = () + + def __new__(cls, index: D, *local_indices: *Ls) -> MultiDimensionIndex[D, *Ls]: + # NOTE: checked at runtime what a checker cannot: a `TypeVarTuple` has no bound. + if not isinstance(index, DimensionIndex) or is_local_dimension(index.dim): + raise TypeError( + f"'MultiDimensionIndex' starts with an index into a non-local dimension, got" + f" '{index!r}'." + ) + if not local_indices: + raise TypeError( + "'MultiDimensionIndex' is a position in a *product*: it needs at least one index" + " into a local dimension." + ) + for local_index in local_indices: + if not isinstance(local_index, LocalDimensionIndex): + raise TypeError( + "'MultiDimensionIndex' continues with indices into local dimensions, got" + f" '{local_index!r}'." + ) + owner = local_index.dim.owner # type: ignore[attr-defined] # a LocalDimensionIndex + if owner is not None and owner.domain is not index.dim: + raise TypeError( + f"'MultiDimensionIndex': '{local_index.dim.__qualname__}' indexes the neighbors" + f" of '{owner.domain.__qualname__}', not of '{index.dim.__qualname__}'." + ) + return super().__new__(cls, (index, *local_indices)) + + def __getnewargs__(self) -> tuple[Any, ...]: + # NOTE: `tuple`'s own passes the elements as one tuple, which `__new__` does not take. + return tuple(cast(tuple[Any, ...], self)) + + @property + def dims(self) -> tuple[Dimension, ...]: + return tuple(index.dim for index in cast(tuple[DimensionIndex, ...], self)) + + def __repr__(self) -> str: + return f"{type(self).__name__}({', '.join(map(repr, cast(tuple[Any, ...], self)))})" diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index d3078fd224..8509e716a6 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -95,7 +95,11 @@ ) -class SparseTag(Tag): ... +@dataclasses.dataclass(frozen=True) +class SparseAxis: + """The offset part that steps along a local dimension of the iterator's own field.""" + + dim: common.Dimension # TODO(havogt): complete implementation and make available for fieldview embedded @@ -184,8 +188,8 @@ def skip_value( # A cartesian shift can be passed as a `common.CartesianConnectivity` tag (carrying its source # axis, codomain, and integer offset); named offsets use a string `Tag` resolved via the offset # provider. -OffsetPart: TypeAlias = Tag | common.CartesianConnectivity | common.IntIndex -CompleteOffset: TypeAlias = tuple[Tag | common.CartesianConnectivity, common.IntIndex] +OffsetPart: TypeAlias = Tag | SparseAxis | common.CartesianConnectivity | common.IntIndex +CompleteOffset: TypeAlias = tuple[Tag | SparseAxis | common.CartesianConnectivity, common.IntIndex] OffsetProviderElem: TypeAlias = common.OffsetProviderElem OffsetProvider: TypeAlias = common.OffsetProvider @@ -194,14 +198,15 @@ def skip_value( IncompleteSparsePositionEntry: TypeAlias = list[Optional[int]] PositionEntry: TypeAlias = SparsePositionEntry | common.IntIndex IncompletePositionEntry: TypeAlias = IncompleteSparsePositionEntry | common.IntIndex -ConcretePosition: TypeAlias = dict[Tag, PositionEntry] -IncompletePosition: TypeAlias = dict[Tag, IncompletePositionEntry] +# NOTE: keyed by the dimension classes themselves; their tags are only how the IR spells them. +ConcretePosition: TypeAlias = dict[common.Dimension, PositionEntry] +IncompletePosition: TypeAlias = dict[common.Dimension, IncompletePositionEntry] Position: TypeAlias = Union[ConcretePosition, IncompletePosition] #: A ``None`` position flags invalid not-a-neighbor results in neighbor-table lookups MaybePosition: TypeAlias = Optional[Position] -NamedFieldIndices: TypeAlias = Mapping[Tag, FieldIndex | SparsePositionEntry] +NamedFieldIndices: TypeAlias = Mapping[common.Dimension, FieldIndex | SparsePositionEntry] @runtime_checkable @@ -433,7 +438,7 @@ def deref(self): return impl -NamedRange: TypeAlias = tuple[Tag | common.Dimension, range] +NamedRange: TypeAlias = tuple[common.Dimension, range] @builtins.cartesian_domain.register(EMBEDDED) @@ -447,12 +452,12 @@ def unstructured_domain(*args: NamedRange) -> runtime.UnstructuredDomain: Domain: TypeAlias = ( - runtime.CartesianDomain | runtime.UnstructuredDomain | dict[str | common.Dimension, range] + runtime.CartesianDomain | runtime.UnstructuredDomain | dict[common.Dimension, range] ) @builtins.named_range.register(EMBEDDED) -def named_range(tag: Tag | common.Dimension, start: int, end: int) -> NamedRange: +def named_range(tag: common.Dimension, start: int, end: int) -> NamedRange: # TODO revisit this pattern after the discussion of 0d-field vs scalar if isinstance(start, ConstantField): start = start.value @@ -522,11 +527,13 @@ def promote_scalars(val: CompositeOfScalarOrField): globals()[math_builtin_name] = decorator(impl) -def _named_range(axis: str, range_: Iterable[int]) -> Iterable[tuple[Tag, common.IntIndex]]: +def _named_range( + axis: common.Dimension, range_: Iterable[int] +) -> Iterable[tuple[common.Dimension, common.IntIndex]]: return ((axis, i) for i in range_) -def _domain_iterator(domain: dict[Tag, range]) -> Iterable[ConcretePosition]: +def _domain_iterator(domain: dict[common.Dimension, range]) -> Iterable[ConcretePosition]: return ( dict(elem) for elem in itertools.product(*(_named_range(axis, rang) for axis, rang in domain.items())) @@ -535,14 +542,14 @@ def _domain_iterator(domain: dict[Tag, range]) -> Iterable[ConcretePosition]: def execute_shift( pos: Position, - tag: Tag | common.CartesianConnectivity, + tag: Tag | SparseAxis | common.CartesianConnectivity, index: common.IntIndex, *, offset_provider: OffsetProvider, ) -> MaybePosition: assert pos is not None - if isinstance(tag, SparseTag): - current_entry = pos[tag] + if isinstance(tag, SparseAxis): + current_entry = pos[tag.dim] assert isinstance(current_entry, list) new_entry = list(current_entry) assert None in new_entry @@ -552,18 +559,18 @@ def execute_shift( for i, p in reversed(list(enumerate(new_entry))): # first shift applies to the last sparse dimensions of that axis type if p is None: - if tag == common.ConstList.tag: + if tag.dim is common.ConstList: new_entry[i] = 0 else: - # NOTE: the sparse tag is the local dimension's; the table over it may be - # keyed by a connectivity sharing it (see `common.connectivity_key_over`). + # NOTE: the table over the local dimension may be keyed by a connectivity + # sharing it (see `common.connectivity_key_over`). offset_implementation = common.get_offset( offset_provider, - common.connectivity_key_over(offset_provider, tag), + common.connectivity_key_over(offset_provider, tag.dim), ) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().domain[0] - cur_index = pos[source_dim.tag] + cur_index = pos[source_dim] assert common.is_int_index(cur_index) if offset_implementation[cur_index, index].as_scalar() in [ None, @@ -574,21 +581,21 @@ def execute_shift( new_entry[i] = index break # the assertions above confirm pos is incomplete casting here to avoid duplicating work in a type guard - return cast(IncompletePosition, pos) | {tag: new_entry} + return cast(IncompletePosition, pos) | {tag.dim: new_entry} if isinstance(tag, common.CartesianConnectivity): new_pos = copy.copy(pos) - value = new_pos.pop(tag.domain_dim.tag) + value = new_pos.pop(tag.domain_dim) assert common.is_int_index(value) - new_pos[tag.codomain.tag] = value + index + tag.offset + new_pos[tag.codomain] = value + index + tag.offset return new_pos offset_implementation = common.get_offset(offset_provider, tag) if common.is_neighbor_table(offset_implementation): source_dim = offset_implementation.__gt_type__().domain[0] - assert source_dim.tag in pos + assert source_dim in pos new_pos = pos.copy() - new_pos.pop(source_dim.tag) - cur_index = pos[source_dim.tag] + new_pos.pop(source_dim) + cur_index = pos[source_dim] assert common.is_int_index(cur_index) if offset_implementation[cur_index, index].as_scalar() in [ None, @@ -598,7 +605,7 @@ def execute_shift( else: new_index = offset_implementation[cur_index, index].as_scalar() assert new_index is not None - new_pos[offset_implementation.codomain.tag] = int(new_index) + new_pos[offset_implementation.codomain] = int(new_index) return new_pos @@ -609,7 +616,7 @@ def _is_list_of_complete_offsets( complete_offsets: list[tuple[Any, Any]], ) -> TypeGuard[list[CompleteOffset]]: return all( - isinstance(tag, (Tag, common.CartesianConnectivity)) + isinstance(tag, (Tag, SparseAxis, common.CartesianConnectivity)) and isinstance(offset, (int, np.integer)) for tag, offset in complete_offsets ) @@ -761,7 +768,7 @@ def _get_axes( def _single_vertical_idx( - indices: NamedFieldIndices, column_axis: Tag, column_index: common.IntIndex + indices: NamedFieldIndices, column_axis: common.Dimension, column_index: common.IntIndex ) -> NamedFieldIndices: transformed = { axis: (index if axis != column_axis else index.start + column_index) # type: ignore[union-attr] # trust me, `index` is range in case of `column_axis` # fmt: off @@ -775,7 +782,7 @@ def _make_tuple( field_or_tuple: tuple[tuple | LocatedField, ...], # arbitrary nesting of tuples of Field named_indices: NamedFieldIndices, *, - column_axis: Tag, + column_axis: common.Dimension, ) -> tuple[tuple | Column, ...]: ... @@ -791,7 +798,7 @@ def _make_tuple( @overload def _make_tuple( - field_or_tuple: LocatedField, named_indices: NamedFieldIndices, *, column_axis: Tag + field_or_tuple: LocatedField, named_indices: NamedFieldIndices, *, column_axis: common.Dimension ) -> Column: ... @@ -808,7 +815,7 @@ def _make_tuple( field_or_tuple: LocatedField | tuple[tuple | LocatedField, ...], named_indices: NamedFieldIndices, *, - column_axis: Optional[Tag] = None, + column_axis: Optional[common.Dimension] = None, ) -> Column | npt.DTypeLike | tuple[tuple | Column | npt.DTypeLike | Undefined, ...] | Undefined: if column_axis is None: if isinstance(field_or_tuple, tuple): @@ -861,7 +868,7 @@ def _make_tuple( class MDIterator: field: LocatedField | tuple[LocatedField | tuple, ...] # arbitrary nesting pos: MaybePosition - column_axis: Optional[Tag] = dataclasses.field(default=None, kw_only=True) + column_axis: Optional[common.Dimension] = dataclasses.field(default=None, kw_only=True) def shift(self, *offsets: OffsetPart) -> MDIterator: complete_offsets = group_offsets(*offsets) @@ -887,9 +894,9 @@ def deref(self) -> Any: axes = _get_axes(self.field, ignore_zero_dims=True) if __debug__: - if not all(axis.tag in shifted_pos.keys() for axis in axes if axis is not None): + if not all(axis in shifted_pos.keys() for axis in axes if axis is not None): raise IndexError("Iterator position doesn't point to valid location for its field.") - slice_column = dict[Tag, range]() + slice_column = dict[common.Dimension, range]() if self.column_axis is not None: column_range = embedded_context.get_closure_column_range() assert column_range is not None @@ -929,18 +936,16 @@ def make_in_iterator( new_pos: Position = pos.copy() for sparse_dim in set(sparse_dimensions): init = [None] * sparse_dimensions.count(sparse_dim) - new_pos[sparse_dim.tag] = init # type: ignore[assignment] # looks like mypy is confused + new_pos[sparse_dim] = init # type: ignore[assignment] # looks like mypy is confused if column_dimension is not None: column_range = embedded_context.get_closure_column_range().unit_range # if we deal with column stencil the column position is just an offset by which the whole column needs to be shifted assert column_range is not None - new_pos[column_dimension.tag] = column_range.start - it = MDIterator( - inp, new_pos, column_axis=column_dimension.tag if column_dimension is not None else None - ) + new_pos[column_dimension] = column_range.start + it = MDIterator(inp, new_pos, column_axis=column_dimension) if len(sparse_dimensions) >= 1: if len(sparse_dimensions) == 1: - return SparseListIterator(it, sparse_dimensions[0].tag) + return SparseListIterator(it, sparse_dimensions[0]) else: raise NotImplementedError( f"More than one local dimension is currently not supported, got {sparse_dimensions}." @@ -966,7 +971,7 @@ def _translate_named_indices( self, _named_indices: NamedFieldIndices ) -> common.AbsoluteIndexSequence: named_indices: Mapping[common.Dimension, FieldIndex | SparsePositionEntry] = { - d: _named_indices[d.tag] for d in self._ndarrayfield.__gt_domain__.dims + d: _named_indices[d] for d in self._ndarrayfield.__gt_domain__.dims } domain_slice: list[common.NamedRange | common.DimensionIndex] = [] for d, v in named_indices.items(): @@ -989,14 +994,14 @@ def field_getitem(self, named_indices: NamedFieldIndices) -> Any: def field_setitem(self, named_indices: NamedFieldIndices, value: Any): if isinstance(self._ndarrayfield, common.MutableField): if isinstance(value, _List): - local_tag = value.local_dim.tag + local_dim = value.local_dim for i, v in enumerate(value): # type:ignore[var-annotated, arg-type] self._ndarrayfield[ - self._translate_named_indices({**named_indices, local_tag: i}) + self._translate_named_indices({**named_indices, local_dim: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ - self._translate_named_indices({**named_indices, common.ConstList.tag: 0}) + self._translate_named_indices({**named_indices, common.ConstList: 0}) ] = value.value else: self._ndarrayfield[self._translate_named_indices(named_indices)] = value @@ -1024,27 +1029,6 @@ def _is_sparse_position_entry( return isinstance(pos, list) -def get_ordered_indices(axes: Iterable[Axis], pos: NamedFieldIndices) -> tuple[FieldIndex, ...]: - res: list[FieldIndex] = [] - sparse_position_tracker: dict[Tag, int] = {} - for axis in axes: - if _is_tuple_axis(axis): - res.append(slice(None)) - else: - assert _is_field_axis(axis) - assert axis.tag in pos - assert isinstance(axis.tag, str) - elem = pos[axis.tag] - if _is_sparse_position_entry(elem): - sparse_position_tracker.setdefault(axis.tag, 0) - res.append(elem[sparse_position_tracker[axis.tag]]) - sparse_position_tracker[axis.tag] += 1 - else: - assert isinstance(elem, (int, np.integer, slice, range)) - res.append(elem) - return tuple(res) - - @overload def _shift_range(range_or_index: range, offset: int) -> slice: ... @@ -1526,19 +1510,19 @@ def sten(*lists): @dataclasses.dataclass(frozen=True) class SparseListIterator: it: ItIterator - list_offset: Tag + list_dim: common.Dimension offsets: Sequence[OffsetPart] = dataclasses.field(default_factory=list, kw_only=True) def deref(self) -> Any: - if self.list_offset == common.ConstList.tag: + if self.list_dim is common.ConstList: return _ConstList( - value=self.it.shift(*self.offsets, SparseTag(self.list_offset), 0).deref() + value=self.it.shift(*self.offsets, SparseAxis(self.list_dim), 0).deref() ) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None - # NOTE: `list_offset` is the local dimension's tag; the table over it may be keyed by a + # NOTE: the table over the local dimension may be keyed by a # connectivity sharing it (see `common.connectivity_key_over`). - connectivity_key = common.connectivity_key_over(offset_provider, self.list_offset) + connectivity_key = common.connectivity_key_over(offset_provider, self.list_dim) connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( @@ -1546,7 +1530,7 @@ def deref(self) -> Any: shifted.deref() for i in range(len(connectivity.domain[1].unit_range)) if ( - shifted := self.it.shift(*self.offsets, SparseTag(self.list_offset), i) + shifted := self.it.shift(*self.offsets, SparseAxis(self.list_dim), i) ).can_deref() ), offset=runtime.Offset(value=connectivity_key), @@ -1556,7 +1540,7 @@ def can_deref(self) -> bool: return self.it.shift(*self.offsets).can_deref() def shift(self, *offsets: OffsetPart) -> SparseListIterator: - return SparseListIterator(self.it, self.list_offset, offsets=[*offsets, *self.offsets]) + return SparseListIterator(self.it, self.list_dim, offsets=[*offsets, *self.offsets]) @dataclasses.dataclass(frozen=True) @@ -1674,10 +1658,17 @@ def impl(*iters: ItIterator): return impl -def _dimension_to_tag( +def _domain_as_dict( domain: runtime.CartesianDomain | runtime.UnstructuredDomain, -) -> dict[Tag, range]: - return {k.tag: v for k, v in domain.items()} +) -> dict[common.Dimension, range]: + result = dict(domain.items()) + for dim in result: + if not isinstance(dim, common.DimensionMeta): + raise TypeError( + f"Domain axis '{dim!r}' is not a dimension; an axis given by its tag can be turned" + " into its dimension with 'gtx.resolve'." + ) + return result def _validate_domain(domain: Domain, table_types: common.TableTypes) -> None: @@ -1754,10 +1745,10 @@ def _extract_column_range(domain) -> common.NamedRange | eve.NothingType: col_range_placeholder.unit_range.is_empty() ) # check it's just the placeholder with empty range column_axis = col_range_placeholder.dim - if column_axis is not None and column_axis.tag in domain: + if column_axis is not None and column_axis in domain: return common.NamedRange( column_axis, - common.UnitRange(domain[column_axis.tag].start, domain[column_axis.tag].stop), + common.UnitRange(domain[column_axis].start, domain[column_axis].stop), ) return eve.NOTHING @@ -1767,13 +1758,13 @@ def _get_output_type( domain_: runtime.CartesianDomain | runtime.UnstructuredDomain, args: tuple[Any, ...], ) -> ts.TypeSpec: - domain = _dimension_to_tag(domain_) + domain = _domain_as_dict(domain_) col_range = _extract_column_range(domain) col_dim: Optional[common.Dimension] = None if isinstance(col_range, common.NamedRange): col_dim = col_range.dim - del domain[col_range.dim.tag] + del domain[col_range.dim] # determine dtype by computing result at one point pos_in_domain = next(iter(_domain_iterator(domain))) @@ -1843,7 +1834,7 @@ def closure( assert embedded_context.within_valid_context() offset_provider = embedded_context.get_offset_provider() _validate_domain(domain_, common.offset_provider_to_type(offset_provider)) - domain: dict[Tag, range] = _dimension_to_tag(domain_) + domain = _domain_as_dict(domain_) if not (isinstance(out, common.Field) or is_tuple_of_field(out)): raise TypeError("'Out' needs to be a located field.") @@ -1852,7 +1843,7 @@ def closure( column_dim = None if isinstance(column_range, common.NamedRange): column_dim = column_range.dim - del domain[column_range.dim.tag] + del domain[column_range.dim] out = as_tuple_field(out) if is_tuple_of_field(out) else _wrap_field(out) promoted_ins = [promote_scalars(inp) for inp in ins] @@ -1868,7 +1859,7 @@ def closure( column_range = cast(common.NamedRange, column_range) col_pos = pos.copy() for k in column_range.unit_range: - col_pos[column_range.dim.tag] = k + col_pos[column_range.dim] = k assert _is_concrete_position(col_pos) out.field_setitem(col_pos, res[k]) # type: ignore[index] diff --git a/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py b/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py index 9e7e73862a..a8694b4118 100644 --- a/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py +++ b/tests/next_tests/unit_tests/iterator_tests/test_embedded_internals.py @@ -151,3 +151,29 @@ def test_func(): def test_lift_accepts_cartesian_dimension_offset(): lifted = embedded.lift(lambda *args: 0)() lifted.shift(common.CartesianConnectivity(K), 1) # must not raise + + +def test_domain_axes_must_be_dimensions(): + from gt4py.next.iterator import embedded as iterator_embedded, runtime as iterator_runtime + + with pytest.raises(TypeError, match="gtx.resolve"): + iterator_embedded._domain_as_dict( + iterator_runtime.CartesianDomain([("pkg.IDim", range(3))]) + ) + + +def test_sparse_axis_steps_along_a_local_dimension(): + """`SparseAxis(dim)` is the offset part of a shift along a field's own local dimension.""" + import gt4py.next as gtx + + from next_tests.toy_connectivity import V2E, V2EDim, Vertex, v2e_arr, v2e_conn + + sparse = gtx.as_field([Vertex, V2EDim], v2e_arr) + with embedded_context.update(offset_provider={V2E.offset_tag: v2e_conn}): + sparse_iterator = embedded.make_in_iterator(sparse, {Vertex: 0}, column_dimension=None) + assert isinstance(sparse_iterator, embedded.SparseListIterator) + # `SparseAxis` fills the sparse entry of the position, which is what derefing a + # `SparseListIterator` does for each neighbor in turn + element = sparse_iterator.it.shift(embedded.SparseAxis(V2EDim), 2) + assert element.deref() == v2e_arr[0][2] + assert tuple(sparse_iterator.deref().values) == tuple(v2e_arr[0]) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index a62c56ecc3..8f5fc07d3c 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -490,9 +490,8 @@ def test_neighbor_index_accepts_numpy_integers(self): assert np.array_equal(V2E[np.int32(1)].asnumpy(), V2E[1].asnumpy()) def test_attribute_errors_are_dsl_errors(self): - from gt4py.next import errors, field_operator + from gt4py.next import Dims, Field, errors from gt4py.next.ffront.func_to_foast import FieldOperatorParser - from gt4py.next import Dims, Field def domain_of(a: Field[Dims[Edge], float]) -> Field[Dims[Vertex], float]: return a(V2E.domain) @@ -791,3 +790,47 @@ def test_layout_order_is_horizontal_local_vertical(self): assert common.order_dimensions([KDim, LsqCoeff, Vertex]) == [Vertex, LsqCoeff, KDim] with pytest.raises(ValueError, match="more than one local dimension"): common.order_dimensions([Vertex, LsqCoeff, V2E.Local]) + + +class TestMultiDimensionIndex: + def test_indexes_a_neighbor_table(self): + table = _table() + position = common.MultiDimensionIndex(Vertex(1), V2E.Local(2)) + assert position.dims == (Vertex, V2E.Local) + assert table[position].as_scalar() == 3 + + def test_is_a_tuple_of_indices(self): + position = common.MultiDimensionIndex(Vertex(1), V2E.Local(2)) + assert position == (Vertex(1), V2E.Local(2)) + assert hash(position) == hash((Vertex(1), V2E.Local(2))) + + def test_pickle_and_copy(self): + import copy + + position = common.MultiDimensionIndex(Vertex(1), V2E.Local(2)) + for clone in ( + pickle.loads(pickle.dumps(position)), + copy.copy(position), + copy.deepcopy(position), + ): + assert type(clone) is common.MultiDimensionIndex + assert clone == position + + @pytest.mark.parametrize( + "indices", + [ + (V2E.Local(0),), + (Vertex(0), Edge(1)), + (Vertex(0), 1), + (0,), + (Vertex(0),), # a position in a product needs a local index + (Edge(0), V2E.Local(1)), # V2E indexes the neighbors of a vertex + ], + ) + def test_rejects_other_shapes(self, indices): + with pytest.raises(TypeError, match="MultiDimensionIndex"): + common.MultiDimensionIndex(*indices) + + def test_accepts_an_owner_less_local_dimension(self): + position = common.MultiDimensionIndex(Vertex(1), LsqCoeff(2)) + assert position.dims == (Vertex, LsqCoeff) diff --git a/typing_tests/test_next.yaml b/typing_tests/test_next.yaml index 3519a3eb82..1d1c834443 100644 --- a/typing_tests/test_next.yaml +++ b/typing_tests/test_next.yaml @@ -373,3 +373,21 @@ reveal_type(first(a)) out: | main:22:17: note: Revealed type is "type[main.V2E.Local]" + + - case: multi_dimension_index + main: | + from gt4py import next as gtx + + class Vertex(gtx.DimensionIndex): ... + class Edge(gtx.DimensionIndex): ... + + class V2E(gtx.NeighborConnectivity[Vertex, Edge]): + class Local(gtx.LocalDimensionIndex): ... + + position = gtx.MultiDimensionIndex(Vertex(3), V2E.Local(1)) + reveal_type(position) + vertex, neighbor = position + reveal_type(neighbor) + out: | + main:10:13: note: Revealed type is "tuple[main.Vertex, main.V2E.Local, fallback=gt4py.next.common.MultiDimensionIndex[main.Vertex, main.V2E.Local]]" + main:12:13: note: Revealed type is "main.V2E.Local"