From 362ac09ce374363cd0600979ca7edb4b863a62f9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 01:03:00 +0200 Subject: [PATCH 1/2] feat[next]: backends support connectivities sharing a local dimension Reductions, sparse arguments and list materialization used to find a connectivity by the local dimension's tag, which only names the owner's table. They now look up a table over the local dimension (common.connectivity_key_over), so a connectivity sharing another one's local dimension works on every backend, bound on its own or together with its owner. DaCe sizes a connectivity array's local dimension from its own table. Lifts the PR 1 xfail markers. --- pyproject.toml | 2 - src/gt4py/next/common.py | 32 +++++++++++++++ src/gt4py/next/embedded/nd_array_field.py | 4 +- src/gt4py/next/iterator/embedded.py | 40 ++++++++++++++----- .../next/iterator/transforms/unroll_reduce.py | 10 +++-- .../codegens/gtfn/gtfn_module.py | 7 +++- .../runners/dace/lowering/gtir_to_sdfg.py | 16 ++++++-- .../dace/lowering/gtir_to_sdfg_lambda.py | 16 ++++++-- .../runners/dace/sdfg_args.py | 21 ++++++++++ tests/next_tests/definitions.py | 39 +----------------- .../test_neighbor_connectivity.py | 38 ++++++++++++++++-- .../test_offset_dimensions_names.py | 24 ++++------- .../transforms_tests/test_unroll_reduce.py | 8 ++-- 13 files changed, 171 insertions(+), 86 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 89a743ce18..3a5ea4af72 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -263,8 +263,6 @@ markers = [ 'uses_ir_if_stmts', 'uses_lift: tests that require backend support for lift builtin function', 'uses_negative_modulo: tests that require backend support for modulo on negative numbers', - 'uses_offset_tag_differing_from_local_dim: tests that shift with an offset tag that differs from the name of its local dimension', - 'uses_offset_tag_differing_from_local_dim_in_reduction: tests that reduce over an offset whose tag differs from the name of its local dimension', 'uses_origin: tests that require backend support for domain origin', 'uses_reduce_with_lambda: tests that use lambdas as reduce functions', 'uses_reduction_with_only_sparse_fields: tests that require backend support for with sparse fields', diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index d91bbff004..083c4d4359 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1471,6 +1471,38 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid get_offset_type: Callable[[OffsetProviderType, str], OffsetProviderTypeElem] = get_offset # type: ignore[assignment] # overload not possible since OffsetProvider and OffsetProviderType overlap +def connectivity_key_over( + offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension +) -> str: + """ + The key of a bound connectivity whose local dimension is `local_dim`. + + Neighbor reductions and sparse fields know only their local dimension, and use its table + for the neighbor count and the skip values. That is the table keyed by the local dimension's + tag, i.e. its owner's, if bound. Otherwise it is a connectivity *sharing* the local dimension + (see `NeighborConnectivity`), keyed by its own tag, which has the same neighbor structure. + + Raises: + KeyError: If no bound connectivity has `local_dim` as its local dimension. + """ + if local_dim.tag in offset_provider: + return local_dim.tag + for key, connectivity in offset_provider.items(): + if isinstance(connectivity, NeighborConnectivityType): + neighbor_dim = connectivity.neighbor_dim + elif is_neighbor_table(connectivity): + neighbor_dim = connectivity.domain.dims[1] + else: + continue + if neighbor_dim is local_dim: + assert isinstance(key, str) + return key + raise KeyError( + f"No connectivity over the local dimension '{local_dim.tag}' is bound in the offset" + f" provider, which has {sorted(map(str, offset_provider))}." + ) + + def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: str) -> bool: """Determine if offset provider has an element for the given offset tag.""" try: diff --git a/src/gt4py/next/embedded/nd_array_field.py b/src/gt4py/next/embedded/nd_array_field.py index f8c422f134..1bbdf3ef82 100644 --- a/src/gt4py/next/embedded/nd_array_field.py +++ b/src/gt4py/next/embedded/nd_array_field.py @@ -988,8 +988,8 @@ def _builtin_op( current_offset_provider = embedded_context.get_offset_provider(None) assert current_offset_provider is not None offset_definition = common.get_offset( - current_offset_provider, axis.tag - ) # assumes offset and local dimension have same name + current_offset_provider, common.connectivity_key_over(current_offset_provider, axis) + ) assert common.is_neighbor_table(offset_definition) new_domain = common.Domain(*[nr for nr in field.domain if nr.dim != axis]) diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 5a9e4011cc..73d7d00231 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -578,7 +578,12 @@ def execute_shift( if tag == _CONST_DIM.tag: new_entry[i] = 0 else: - offset_implementation = common.get_offset(offset_provider, tag) + # NOTE: the sparse tag is the local dimension's; the table over it may be + # keyed by a connectivity sharing it (see `common.connectivity_key_over`). + offset_implementation = common.get_offset( + offset_provider, + common.connectivity_key_over(offset_provider, common.resolve(tag)), + ) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().source_dim cur_index = pos[source_dim.tag] @@ -1009,7 +1014,7 @@ def field_setitem(self, named_indices: NamedFieldIndices, value: Any): if isinstance(value, _List): for i, v in enumerate(value): # type:ignore[var-annotated, arg-type] self._ndarrayfield[ - self._translate_named_indices({**named_indices, value.offset.value: i}) # type: ignore[dict-item] + self._translate_named_indices({**named_indices, value.local_dim.tag: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ @@ -1420,16 +1425,26 @@ def __getitem__(self, i: int): return self.values[i] def __gt_type__(self) -> ts.ListType: - offset_tag = self.offset.value - assert isinstance(offset_tag, str) element_type = type_translation.from_value(self.values[0]) assert isinstance(element_type, ts.DataType) + return ts.ListType(element_type=element_type, offset_type=self.local_dim) + + @property + def local_dim(self) -> common.Dimension: + """ + The local dimension the list runs along. + + The neighbor dimension of the connectivity the list was built with, which is not + necessarily named like its offset: a connectivity can share another one's local + dimension. + """ + offset_tag = self.offset.value + assert isinstance(offset_tag, str) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None connectivity = common.get_offset(offset_provider, offset_tag) assert common.is_neighbor_table(connectivity) - local_dim = connectivity.__gt_type__().neighbor_dim - return ts.ListType(element_type=element_type, offset_type=local_dim) + return connectivity.__gt_type__().neighbor_dim @dataclasses.dataclass(frozen=True) @@ -1543,7 +1558,12 @@ def deref(self) -> Any: ) offset_provider = embedded_context.get_offset_provider() assert offset_provider is not None - connectivity = common.get_offset(offset_provider, self.list_offset) + # NOTE: `list_offset` is the local dimension's tag; the table over it may be keyed by a + # connectivity sharing it (see `common.connectivity_key_over`). + connectivity_key = common.connectivity_key_over( + offset_provider, common.resolve(self.list_offset) + ) + connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( values=tuple( @@ -1553,7 +1573,7 @@ def deref(self) -> Any: shifted := self.it.shift(*self.offsets, SparseTag(self.list_offset), i) ).can_deref() ), - offset=runtime.Offset(value=self.list_offset), + offset=runtime.Offset(value=connectivity_key), ) def can_deref(self) -> bool: @@ -1800,7 +1820,9 @@ def _fieldspec_list_to_value( offset_provider = embedded_context.get_offset_provider() offset_type = type_.offset_type assert isinstance(offset_type, common.DimensionMeta) - connectivity = common.get_offset(offset_provider, offset_type.tag) + connectivity = common.get_offset( + offset_provider, common.connectivity_key_over(offset_provider, offset_type) + ) assert common.is_neighbor_table(connectivity) return domain.insert( len(domain), diff --git a/src/gt4py/next/iterator/transforms/unroll_reduce.py b/src/gt4py/next/iterator/transforms/unroll_reduce.py index cfb7bb3226..22491346bb 100644 --- a/src/gt4py/next/iterator/transforms/unroll_reduce.py +++ b/src/gt4py/next/iterator/transforms/unroll_reduce.py @@ -40,11 +40,11 @@ def _get_neighbors_args(reduce_args: Iterable[itir.Expr]) -> Iterator[itir.FunCa return filter(_is_neighbors_or_lifted_and_neighbors, flat_reduce_args) -def _get_partial_offset_tags(reduce_args: Iterable[itir.Expr]) -> Iterable[str]: +def _get_partial_local_dims(reduce_args: Iterable[itir.Expr]) -> Iterable[common.Dimension]: assert all(isinstance(arg.type, ts.ListType) for arg in reduce_args) return [ - arg.type.offset_type.tag # type: ignore[union-attr] # checked in previous lines + arg.type.offset_type # type: ignore[union-attr] # checked in previous lines for arg in reduce_args if arg.type.offset_type is not None # type: ignore[union-attr] # checked in previous lines ] @@ -59,8 +59,10 @@ def _get_connectivity( raise ValueError("Expected a call to a 'reduce' object, i.e. 'reduce(...)(...)'.") connectivities: list[common.NeighborConnectivityType] = [] - for o in _get_partial_offset_tags(applied_reduce_node.args): - conn = common.get_offset_type(offset_provider_type, o) + for local_dim in _get_partial_local_dims(applied_reduce_node.args): + conn = common.get_offset_type( + offset_provider_type, common.connectivity_key_over(offset_provider_type, local_dim) + ) assert isinstance(conn, common.NeighborConnectivityType) connectivities.append(conn) diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index ca67fb18af..36a80cb97c 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -91,7 +91,12 @@ def _process_regular_arguments( # NOTE: the tag is the offset-provider key, and its mangled form names the # `generated::_t` tag type. A legacy `FieldOffset` carries it as `value`. dim_name = dim.value if isinstance(dim, fbuiltins.FieldOffset) else dim.tag - connectivity = common.get_offset_type(offset_provider_type, dim_name) + connectivity = common.get_offset_type( + offset_provider_type, + dim_name + if isinstance(dim, fbuiltins.FieldOffset) + else common.connectivity_key_over(offset_provider_type, dim), + ) assert isinstance(connectivity, common.NeighborConnectivityType) size = connectivity.max_neighbors arg = f"gridtools::sid::dimension_to_tuple_like({arg})" diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py index c4b1526bdc..d3026b3ea2 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg.py @@ -89,6 +89,11 @@ class DataflowBuilder(Protocol): @abc.abstractmethod def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: ... + @abc.abstractmethod + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + """The offset of a connectivity over `local_dim`, see `common.connectivity_key_over`.""" + ... + @abc.abstractmethod def unique_nsdfg_name(self, prefix: str) -> str: ... @@ -564,6 +569,9 @@ class GTIRToSDFG(eve.NodeVisitor, SDFGBuilder): def get_offset_provider_type(self, offset: str) -> gtx_common.OffsetProviderTypeElem: return gtx_common.get_offset_type(self.offset_provider_type, offset) + def connectivity_key_over(self, local_dim: gtx_common.Dimension) -> str: + return gtx_common.connectivity_key_over(self.offset_provider_type, local_dim) + def make_field( self, data_node: dace_nodes.AccessNode, @@ -578,10 +586,12 @@ def make_field( # the local dimension is converted into `ListType` data element if not isinstance(data_type.dtype, ts.ScalarType): raise ValueError(f"Invalid field type {data_type}.") - if not gtx_common.has_offset(self.offset_provider_type, local_dim.tag): + try: + self.connectivity_key_over(local_dim) + except KeyError as ex: raise ValueError( f"The provided local dimension {local_dim} does not match any offset provider type." - ) + ) from ex local_type = ts.ListType(element_type=data_type.dtype, offset_type=local_dim) field_type = ts.FieldType( dims=[dim for dim in data_type.dims if dim != local_dim], dtype=local_type @@ -839,7 +849,7 @@ def _make_array_shape_and_strides( for dim in dims: if dim.kind == gtx_common.DimensionKind.LOCAL: # for local dimension, the size is taken from the associated connectivity type - shape.append(neighbor_table_types[dim.tag].max_neighbors) + shape.append(gtx_dace_args.local_dimension_size(name, dim, neighbor_table_types)) elif gtx_dace_args.is_connectivity_identifier(name, self.offset_provider_type): # we use symbolic size for the global dimension of a connectivity shape.append(gtx_dace_args.field_size_symbol(name, dim, neighbor_table_types)) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 924758c727..2ac65d0e58 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 @@ -1321,7 +1321,9 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if offset_type == _CONST_DIM: # this input argument is the result of `make_const_list` continue - offset_provider_t = self.subgraph_builder.get_offset_provider_type(offset_type.tag) + offset_provider_t = self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) input_conn_types[offset_type] = offset_provider_t @@ -1375,7 +1377,9 @@ def _visit_map(self, node: gtir.FunCall) -> ValueExpr: if conn_type.has_skip_values: # In case the `map_list` input expressions contain skip values, we use # the connectivity-based offset provider as mask for map computation. - conn_data = gtx_dace_args.connectivity_identifier(offset_type.tag) + conn_data = gtx_dace_args.connectivity_identifier( + self.subgraph_builder.connectivity_key_over(offset_type) + ) conn_desc = self.sdfg.arrays[conn_data] conn_desc.transient = False @@ -1477,7 +1481,9 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: and input_expr.gt_dtype.offset_type is not None ) offset_type = input_expr.gt_dtype.offset_type - offset_provider_type = self.subgraph_builder.get_offset_provider_type(offset_type.tag) + offset_provider_type = self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) inp_conn = "_in" @@ -1489,7 +1495,9 @@ def _visit_reduce(self, node: gtir.FunCall) -> ValueExpr: and input_expr.gt_dtype.offset_type is not None ) offset_type = input_expr.gt_dtype.offset_type - connectivity = gtx_dace_args.connectivity_identifier(offset_type.tag) + connectivity = gtx_dace_args.connectivity_identifier( + self.subgraph_builder.connectivity_key_over(offset_type) + ) self.sdfg.arrays[connectivity].transient = False reduce_node = gtx_library_nodes.ReduceWithSkipValues( diff --git a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py index 5b0b0601fc..c614719aff 100644 --- a/src/gt4py/next/program_processors/runners/dace/sdfg_args.py +++ b/src/gt4py/next/program_processors/runners/dace/sdfg_args.py @@ -111,6 +111,27 @@ def field_stride_symbol( return _field_symbol(field_name, dim, "stride", offset_provider_type) +def local_dimension_size( + field_name: str, + dim: gtx_common.Dimension, + neighbor_table_types: dict[str, gtx_common.NeighborConnectivityType], +) -> int: + """ + Number of neighbors along the local dimension `dim` of the field or connectivity table. + + A connectivity table has its own neighbor count. Any other field finds it in a table over + `dim`: normally the one keyed by `dim`'s tag, but a connectivity sharing `dim` with another + one (see `NeighborConnectivity`) is keyed by its own tag, and may be the only one bound. + """ + if (m := CONNECTIVITY_INDENTIFIER_RE.match(field_name)) is not None: + own_type = neighbor_table_types[gtx_common.from_codegen_name(m[1])] + if own_type.neighbor_dim == dim: + return own_type.max_neighbors + return neighbor_table_types[ + gtx_common.connectivity_key_over(neighbor_table_types, dim) + ].max_neighbors + + def _range_symbol_name(field_name: str, dim: gtx_common.Dimension) -> str: """Common part of the name for the range start/stop symbols.""" field_range = im.call("get_domain_range")(field_name, dim) diff --git a/tests/next_tests/definitions.py b/tests/next_tests/definitions.py index 86a4229860..eb659525ed 100644 --- a/tests/next_tests/definitions.py +++ b/tests/next_tests/definitions.py @@ -98,10 +98,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): USES_INDEX_FIELDS = "uses_index_fields" USES_LIFT = "uses_lift" USES_NEGATIVE_MODULO = "uses_negative_modulo" -USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM = "uses_offset_tag_differing_from_local_dim" -USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION = ( - "uses_offset_tag_differing_from_local_dim_in_reduction" -) USES_ORIGIN = "uses_origin" USES_REDUCE_WITH_LAMBDA = "uses_reduce_with_lambda" USES_SCAN = "uses_scan" @@ -142,10 +138,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): #: the connectivity is looked up in the offset provider by the *local dimension's* name. #: Lifted for the gtfn shift path by #1789; see #: `regression_tests/ffront_tests/test_offset_dimensions_names.py`. -OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE = ( - "'{marker}': '{backend}' looks the connectivity up by the local dimension's name," - " so it must equal the offset tag" -) # Index-only vs. consequential markers: # A `uses_*` marker only affects execution if it appears in one of the skip lists below (and thus # in `BACKEND_SKIP_TEST_MATRIX`); such a marker is "consequential" -- it applies the listed @@ -181,16 +173,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_SCAN_IN_STENCIL, XFAIL, BINDINGS_UNSUPPORTED_MESSAGE), (USES_SPARSE_FIELDS, XFAIL, UNSUPPORTED_MESSAGE), (USES_TUPLE_ITERATOR, XFAIL, UNSUPPORTED_MESSAGE), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] ) EMBEDDED_SKIP_LIST = [ @@ -202,11 +184,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): ), # we can't extract the field type from scan args (EMBEDDED_CONCAT_WHERE_INFINITE_DOMAIN, XFAIL, UNSUPPORTED_MESSAGE), (EMBEDDED_CONCAT_WHERE_NON_CONTIGUOUS_DOMAIN, XFAIL, UNSUPPORTED_MESSAGE), - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] JAX_EMBEDDED_SKIP_LIST = EMBEDDED_SKIP_LIST + [ (USES_PROGRAM_WITH_SLICED_OUT_ARGUMENTS, XFAIL, UNSUPPORTED_MESSAGE), @@ -217,15 +194,7 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_TUPLES_ARGS_WITH_DIFFERENT_BUT_PROMOTABLE_DIMS, XFAIL, UNSUPPORTED_MESSAGE), (USES_CONCAT_WHERE, XFAIL, UNSUPPORTED_MESSAGE), ] -GTIR_EMBEDDED_SKIP_LIST = ROUNDTRIP_SKIP_LIST + [ - # NOTE: not in `ROUNDTRIP_SKIP_LIST`: the roundtrip backend passes this, only the - # lower-level `iterator/embedded.py` execution keys on the local dimension's name. - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), -] +GTIR_EMBEDDED_SKIP_LIST = ROUNDTRIP_SKIP_LIST GTFN_SKIP_TEST_LIST = ( COMMON_SKIP_TEST_LIST + DOMAIN_INFERENCE_SKIP_LIST @@ -236,12 +205,6 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum): (USES_STRIDED_NEIGHBOR_OFFSET, XFAIL, BINDINGS_UNSUPPORTED_MESSAGE), # max_over broken, see https://github.com/GridTools/gt4py/issues/1289 (USES_MAX_OVER, XFAIL, UNSUPPORTED_MESSAGE), - # NOTE: only the reduction; #1789 lifted this for the shift path. - ( - USES_OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_IN_REDUCTION, - XFAIL, - OFFSET_TAG_DIFFERING_FROM_LOCAL_DIM_MESSAGE, - ), ] ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_neighbor_connectivity.py index 1b92138e0e..9858fc0bb0 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 @@ -8,6 +8,7 @@ """A `NeighborConnectivity` declaration used directly in DSL code, on every backend.""" +import dataclasses import typing import numpy as np @@ -120,7 +121,6 @@ def testee(a: Field[Dims[E], float], out: Field[Dims[V], float]): @pytest.mark.uses_unstructured_shift -@pytest.mark.uses_offset_tag_differing_from_local_dim def test_shift_through_a_shared_local_dimension(case): @gtx.field_operator def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: @@ -130,8 +130,6 @@ def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: @pytest.mark.uses_unstructured_shift -@pytest.mark.uses_offset_tag_differing_from_local_dim -@pytest.mark.uses_offset_tag_differing_from_local_dim_in_reduction def test_reduction_through_a_shared_local_dimension(case): @gtx.field_operator def testee( @@ -145,3 +143,37 @@ def testee( testee, lambda s, a: np.sum(s * a[_table(case, V2EShared)] - a[_table(case)], axis=1), ) + + +@pytest.fixture +def case_without_owner(case): + """Only the sharing connectivity is bound: enough for a shift, which needs only its table.""" + return dataclasses.replace( + case, offset_provider={V2EShared.offset_tag: case.offset_provider[V2EShared.offset_tag]} + ) + + +@pytest.mark.uses_unstructured_shift +def test_shift_through_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee(a: Field[Dims[E], float]) -> Field[Dims[V], float]: + return a(V2EShared[1]) + + cases.verify_with_default_data( + case_without_owner, testee, lambda a: a[_table(case_without_owner, V2EShared)[:, 1]] + ) + + +@pytest.mark.uses_unstructured_shift +def test_reduction_through_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee( + s: Field[Dims[V, V2E.Local], float], a: Field[Dims[E], float] + ) -> Field[Dims[V], float]: + return neighbor_sum(s * a(V2EShared), axis=V2E.Local) + + cases.verify_with_default_data( + case_without_owner, + testee, + lambda s, a: np.sum(s * a[_table(case_without_owner, V2EShared)], axis=1), + ) 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 3cf6519d27..df066e21cf 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 @@ -56,16 +56,8 @@ class Neigh(gtx.LocalDimensionIndex): ... OffB = gtx.FieldOffset("OffB", source=E, target=(V, Neigh)) -def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases.Case: - """ - A `Case` whose offset provider holds exactly one connectivity, keyed on its tag. - - One entry per `Case` on purpose: DaCe walks every provider entry while building the - SDFG, and looks a connectivity up by its *local dimension's* name - (`gtir_to_sdfg.py`, constraint A4). A second, non-conforming entry would therefore - fail a program that does not even use it, and the cell under test would be measuring - the wrong thing. - """ +def _case(exec_alloc_descriptor, tags: tuple[str, ...], local_dim: common.Dimension) -> cases.Case: + """A `Case` binding the same table under each of `tags`.""" 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. @@ -84,6 +76,7 @@ def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases 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, @@ -93,12 +86,14 @@ def _case(exec_alloc_descriptor, tag: str, local_dim: common.Dimension) -> cases @pytest.fixture def case_tag_vs_variable_name(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, TaggedOffDim.tag, TaggedOffDim) + return _case(exec_alloc_descriptor, (TaggedOffDim.tag,), TaggedOffDim) @pytest.fixture def case_tag_vs_local_dim(exec_alloc_descriptor): - return _case(exec_alloc_descriptor, "OffB", Neigh) + # 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: @@ -135,11 +130,9 @@ def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: # --- N3: the tag differs from the local dimension's name -------------------------- -# Lifted for the gtfn shift path by #1789; still required elsewhere, which is what -# the markers below record. +# The shape of a connectivity sharing another one's local dimension. -@pytest.mark.uses_offset_tag_differing_from_local_dim def test_shift_tag_differs_from_local_dim_name(case_tag_vs_local_dim): """ Ensure a shift works with an offset tag that differs from the local dimension's name. @@ -160,7 +153,6 @@ def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: ) -@pytest.mark.uses_offset_tag_differing_from_local_dim_in_reduction def test_reduction_tag_differs_from_local_dim_name(case_tag_vs_local_dim): @gtx.field_operator def foo(a: Field[Dims[E], float]) -> Field[Dims[V], float]: 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 3f93017c67..7c81aa5514 100644 --- a/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py +++ b/tests/next_tests/unit_tests/iterator_tests/transforms_tests/test_unroll_reduce.py @@ -11,7 +11,7 @@ from gt4py.next import common, utils from gt4py.next.iterator import ir from gt4py.next.iterator.ir_utils import ir_makers as im -from gt4py.next.iterator.transforms.unroll_reduce import UnrollReduce, _get_partial_offset_tags +from gt4py.next.iterator.transforms.unroll_reduce import UnrollReduce, _get_partial_local_dims from gt4py.next.type_system import type_specifications as ts @@ -98,10 +98,10 @@ def reduction_if(): "reduction_if", ], ) -def test_get_partial_offsets(reduction, request): - partial_offsets = _get_partial_offset_tags(request.getfixturevalue(reduction).args) +def test_get_partial_local_dims(reduction, request): + partial_local_dims = _get_partial_local_dims(request.getfixturevalue(reduction).args) - assert set(partial_offsets) == {Dim.tag} + assert set(partial_local_dims) == {Dim} def _expected(red, max_neighbors, has_skip_values, shifted_arg=0): From 22b8382c878efdd51b990ca0d8786094f82dec19 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Tue, 22 Sep 2026 03:19:55 +0200 Subject: [PATCH 2/2] fix[next]: address review of shared local dimensions in backends - iterator tracing and embedded shift name a sharing connectivity by offset_tag - DaCe if/scan/concat_where/const-list sites find the table over the local dimension - connectivity_key_over takes a tag, avoids resolve, and picks deterministically - embedded map_list compares lists by local dimension, not by offset - a sharer must have its owner's origin; counts are compared one by one - ADR 0029: sharers must have the owner's neighbor structure --- .../ADRs/next/0029-Connectivities_As_Types.md | 23 +++--- src/gt4py/next/common.py | 71 ++++++++++++------- src/gt4py/next/iterator/embedded.py | 23 +++--- src/gt4py/next/iterator/tracing.py | 2 +- .../lowering/gtir_to_sdfg_concat_where.py | 4 +- .../dace/lowering/gtir_to_sdfg_lambda.py | 6 +- .../dace/lowering/gtir_to_sdfg_primitives.py | 4 +- .../dace/lowering/gtir_to_sdfg_scan.py | 4 +- .../test_neighbor_connectivity.py | 20 ++++++ .../test_with_toy_connectivity.py | 49 +++++++++++++ .../dace_tests/test_dace_utils.py | 31 ++++++++ .../unit_tests/test_neighbor_connectivity.py | 58 ++++++++++++++- 12 files changed, 242 insertions(+), 53 deletions(-) diff --git a/docs/development/ADRs/next/0029-Connectivities_As_Types.md b/docs/development/ADRs/next/0029-Connectivities_As_Types.md index 150976307b..e2b433fcd6 100644 --- a/docs/development/ADRs/next/0029-Connectivities_As_Types.md +++ b/docs/development/ADRs/next/0029-Connectivities_As_Types.md @@ -124,15 +124,20 @@ whose tag is the connectivity's `offset_tag`: 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. + Shifts find the table by that tag. Reductions and sparse arguments know only + the local dimension, and take its neighbor count and skip values from a table + over it (`common.connectivity_key_over`): the owner's if bound, else the + sharer with the smallest tag. Connectivities sharing a local dimension must + therefore have the same neighbor *structure* — the same count, and a skip value + at the same positions — which is what sharing a neighbor axis means; + `check_offset_provider` enforces it for the tables it is given. + +`V2E.Local` inside DSL code types as that local dimension, 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 diff --git a/src/gt4py/next/common.py b/src/gt4py/next/common.py index 083c4d4359..5ffc813f8c 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1472,35 +1472,45 @@ def get_offset(offset_provider: OffsetProvider, offset_tag: str) -> OffsetProvid def connectivity_key_over( - offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension + offset_provider: OffsetProvider | OffsetProviderType, local_dim: Dimension | Tag ) -> str: """ - The key of a bound connectivity whose local dimension is `local_dim`. + The key of a bound connectivity whose local dimension is `local_dim` (a dimension or its tag). Neighbor reductions and sparse fields know only their local dimension, and use its table for the neighbor count and the skip values. That is the table keyed by the local dimension's - tag, i.e. its owner's, if bound. Otherwise it is a connectivity *sharing* the local dimension - (see `NeighborConnectivity`), keyed by its own tag, which has the same neighbor structure. + tag, i.e. its owner's, if bound. Otherwise it is one of the connectivities *sharing* the local + dimension (see `NeighborConnectivity`), each keyed by its own tag; the smallest key is taken, + so the choice does not depend on the order of the provider. Connectivities sharing a local + dimension have the same neighbor structure (see `check_offset_provider`), so which one does + not matter. Raises: KeyError: If no bound connectivity has `local_dim` as its local dimension. """ - if local_dim.tag in offset_provider: - return local_dim.tag - for key, connectivity in offset_provider.items(): - if isinstance(connectivity, NeighborConnectivityType): - neighbor_dim = connectivity.neighbor_dim - elif is_neighbor_table(connectivity): - neighbor_dim = connectivity.domain.dims[1] - else: - continue - if neighbor_dim is local_dim: - assert isinstance(key, str) - return key - raise KeyError( - f"No connectivity over the local dimension '{local_dim.tag}' is bound in the offset" - f" provider, which has {sorted(map(str, offset_provider))}." - ) + local_tag = local_dim if isinstance(local_dim, str) else local_dim.tag + if local_tag in offset_provider: + return local_tag + candidates = [ + key + for key, connectivity in offset_provider.items() + if (neighbor_dim := _neighbor_dim_of(connectivity)) is not None + and neighbor_dim.tag == local_tag + ] + if not candidates: + raise KeyError( + f"No connectivity over the local dimension '{local_tag}' is bound in the offset" + f" provider, which has {sorted(map(str, offset_provider))}." + ) + return min(candidates) + + +def _neighbor_dim_of(connectivity: Any) -> Optional[Dimension]: + if isinstance(connectivity, NeighborConnectivityType): + return connectivity.neighbor_dim + if is_neighbor_table(connectivity): + return connectivity.domain.dims[1] + return None def has_offset(offset_provider: OffsetProvider | OffsetProviderType, offset_tag: str) -> bool: @@ -2160,14 +2170,23 @@ def __init_subclass__( if local.owner is not None and local.owner.tag != cls.tag: # Sharing another connectivity's local dimension: the neighbor structure is the # owner's, including its counts, and the sharing connectivity is named by its own tag. - if (max_neighbors, min_neighbors) != (None, None) and ( - max_neighbors, - min_neighbors, - ) != (local.max_neighbors, local.min_neighbors): + owner_name = local.owner.__qualname__ + if origin is not local.owner.origin: raise TypeError( - f"'{name}' shares the local dimension of '{local.owner.__qualname__}', whose" - " neighbor counts are declared by its owner." + f"'{name}' cannot share the local dimension of '{owner_name}': it has origin" + f" '{origin}', but the neighbors of '{owner_name}' are those of" + f" '{local.owner.origin}'." ) + for count_name, count in ( + ("max_neighbors", max_neighbors), + ("min_neighbors", min_neighbors), + ): + if count is not None and count != getattr(local, count_name): + raise TypeError( + f"'{name}': '{count_name}={count}' contradicts the local dimension it" + f" shares with '{owner_name}', which declares" + f" {count_name}={getattr(local, count_name)}." + ) cls.origin, cls.codomain = origin, codomain return for count_name, count in ( diff --git a/src/gt4py/next/iterator/embedded.py b/src/gt4py/next/iterator/embedded.py index 73d7d00231..7062053223 100644 --- a/src/gt4py/next/iterator/embedded.py +++ b/src/gt4py/next/iterator/embedded.py @@ -582,7 +582,7 @@ def execute_shift( # keyed by a connectivity sharing it (see `common.connectivity_key_over`). offset_implementation = common.get_offset( offset_provider, - common.connectivity_key_over(offset_provider, common.resolve(tag)), + common.connectivity_key_over(offset_provider, tag), ) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().source_dim @@ -1012,9 +1012,10 @@ def field_getitem(self, named_indices: NamedFieldIndices) -> Any: def field_setitem(self, named_indices: NamedFieldIndices, value: Any): if isinstance(self._ndarrayfield, common.MutableField): if isinstance(value, _List): + local_tag = value.local_dim.tag for i, v in enumerate(value): # type:ignore[var-annotated, arg-type] self._ndarrayfield[ - self._translate_named_indices({**named_indices, value.local_dim.tag: i}) + self._translate_named_indices({**named_indices, local_tag: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ @@ -1467,7 +1468,7 @@ def _as_offset_tag( offset: runtime.Offset | type[common.NeighborConnectivity] | OffsetPart, ) -> OffsetPart: if isinstance(offset, common.ConnectivityMeta): - return common.local_dimension_of(offset).tag + return offset.offset_tag return offset.value if isinstance(offset, runtime.Offset) else offset @@ -1500,12 +1501,14 @@ def list_get(i, lst: _List[Optional[DT]]) -> Optional[DT] | Undefined: def _get_offset(*lists: _List | _ConstList) -> Optional[runtime.Offset]: - offsets = set((lst.offset for lst in lists if hasattr(lst, "offset"))) - if len(offsets) == 0: + neighbor_lists = [lst for lst in lists if isinstance(lst, _List)] + if len(neighbor_lists) == 0: return None - if len(offsets) == 1: - return offsets.pop() - raise AssertionError("All lists must have the same offset.") + # NOTE: compared by local dimension, not by offset: a connectivity and one sharing its local + # dimension build lists along the same axis. + if len({lst.local_dim for lst in neighbor_lists}) != 1: + raise AssertionError("All lists must run along the same local dimension.") + return neighbor_lists[0].offset @builtins.map_list.register(EMBEDDED) @@ -1560,9 +1563,7 @@ def deref(self) -> Any: assert offset_provider is not None # NOTE: `list_offset` is the local dimension's tag; the table over it may be keyed by a # connectivity sharing it (see `common.connectivity_key_over`). - connectivity_key = common.connectivity_key_over( - offset_provider, common.resolve(self.list_offset) - ) + connectivity_key = common.connectivity_key_over(offset_provider, self.list_offset) connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( diff --git a/src/gt4py/next/iterator/tracing.py b/src/gt4py/next/iterator/tracing.py index 094c85925f..ec8db888b1 100644 --- a/src/gt4py/next/iterator/tracing.py +++ b/src/gt4py/next/iterator/tracing.py @@ -154,7 +154,7 @@ def make_node(o): assert o.offset == 0 return im.cartesian_offset(o.domain_dim, o.codomain) if isinstance(o, common.ConnectivityMeta): - return OffsetLiteral(value=common.local_dimension_of(o).tag) + return OffsetLiteral(value=o.offset_tag) if callable(o): if o.__name__ == "": return lambdadef(o) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py index 737648b911..60827eeedc 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_concat_where.py @@ -254,7 +254,9 @@ def translate_concat_where( local_dim = node.type.dtype.offset_type assert local_dim is not None dtype = gtx_dace_args.as_dace_type(node.type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type(local_dim.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(local_dim) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) local_idx = gtx_common.order_dimensions([*output_dims, local_dim]).index(local_dim) output_shape.insert(local_idx, offset_provider_type.max_neighbors) diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_lambda.py index 2ac65d0e58..3b44ecc4c2 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 @@ -769,7 +769,9 @@ def _visit_if_branch_arg( local_dim = arg.gt_dtype.offset_type assert local_dim is not None assert isinstance( - self.subgraph_builder.get_offset_provider_type(local_dim.tag), + self.subgraph_builder.get_offset_provider_type( + self.subgraph_builder.connectivity_key_over(local_dim) + ), gtx_common.NeighborConnectivityType, ) # find position of the local dimension in the field layout @@ -1441,7 +1443,7 @@ def _broadcast_const_list( ) -> ValueExpr: assert list_type.offset_type is not None offset_provider_t = self.subgraph_builder.get_offset_provider_type( - list_type.offset_type.tag + self.subgraph_builder.connectivity_key_over(list_type.offset_type) ) assert isinstance(offset_provider_t, gtx_common.NeighborConnectivityType) local_size = offset_provider_t.max_neighbors diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py index e265377a45..50a445146c 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_primitives.py @@ -325,7 +325,9 @@ def _construct_if_branch_output( assert out_type.dtype.offset_type is not None assert isinstance(out_type.dtype.element_type, ts.ScalarType) dtype = gtx_dace_args.as_dace_type(out_type.dtype.element_type) - offset_provider_type = sdfg_builder.get_offset_provider_type(out_type.dtype.offset_type.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(out_type.dtype.offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) shape = [*shape, offset_provider_type.max_neighbors] diff --git a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py index da75325989..fb3099aa3d 100644 --- a/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py +++ b/src/gt4py/next/program_processors/runners/dace/lowering/gtir_to_sdfg_scan.py @@ -384,7 +384,9 @@ def get_scan_output_shape( assert isinstance(scan_init_data.gt_type, ts.ListType) assert scan_init_data.gt_type.offset_type offset_type = scan_init_data.gt_type.offset_type - offset_provider_type = sdfg_builder.get_offset_provider_type(offset_type.tag) + offset_provider_type = sdfg_builder.get_offset_provider_type( + sdfg_builder.connectivity_key_over(offset_type) + ) assert isinstance(offset_provider_type, gtx_common.NeighborConnectivityType) list_size = offset_provider_type.max_neighbors return [scan_column_size, dace.symbolic.SymExpr(list_size)] diff --git a/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 9858fc0bb0..bce54f4e6a 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 @@ -177,3 +177,23 @@ def testee( testee, lambda s, a: np.sum(s * a[_table(case_without_owner, V2EShared)], axis=1), ) + + +@pytest.mark.uses_unstructured_shift +@pytest.mark.uses_if_stmts +def test_if_over_a_shared_local_dimension_without_its_owner(case_without_owner): + @gtx.field_operator + def testee(a: Field[Dims[E], float], flag: bool) -> Field[Dims[V], float]: + if flag: + s = a(V2EShared) + else: + s = a(V2EShared) * 2.0 + return neighbor_sum(s, axis=V2E.Local) + + cases.verify_with_default_data( + case_without_owner, + testee, + lambda a, flag: ( + np.sum(a[_table(case_without_owner, V2EShared)], axis=1) * (1.0 if flag else 2.0) + ), + ) diff --git a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py index fdf8cf5114..d055f1f094 100644 --- a/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py +++ b/tests/next_tests/integration_tests/multi_feature_tests/iterator_tests/test_with_toy_connectivity.py @@ -434,3 +434,52 @@ def test_sparse_shifted_stencil_reduce(program_processor): if validate: assert np.allclose(out.asnumpy(), ref) + + +class V2EShared(gtx.NeighborConnectivity[Vertex, Edge]): + """Shares `V2E`'s local dimension; bound to the table with its columns reversed.""" + + Local = V2EDim + + +v2e_shared_arr = np.ascontiguousarray(v2e_arr[:, ::-1]) +v2e_shared_conn = gtx.as_connectivity( + domain={Vertex: v2e_shared_arr.shape[0], V2EDim: v2e_shared_arr.shape[1]}, + codomain=Edge, + data=v2e_shared_arr, +) + + +@fundef +def shift_through_sharer(in_edges): + return deref(shift(V2EShared, 1)(in_edges)) + + +@fundef +def owner_times_sharer(in_edges): + return reduce(plus, 0)( + map_list(multiplies)(neighbors(V2EShared, in_edges), neighbors(V2E, in_edges)) + ) + + +@pytest.mark.parametrize( + "stencil, ref", + [ + (shift_through_sharer, v2e_shared_arr[:, 1]), + (owner_times_sharer, np.sum(v2e_shared_arr * v2e_arr, axis=1)), + ], +) +def test_connectivity_sharing_a_local_dimension(program_processor, stencil, ref): + program_processor, validate = program_processor + inp = edge_index_field() + out = gtx.as_field([Vertex], np.zeros([9], dtype=inp.dtype)) + + run_processor( + stencil[{Vertex: range(0, 9)}], + program_processor, + inp, + out=out, + offset_provider={V2E.offset_tag: v2e_conn, V2EShared.offset_tag: v2e_shared_conn}, + ) + if validate: + assert np.allclose(out.asnumpy(), ref) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py index 6fdec5a21f..78f1045eb9 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_utils.py @@ -18,3 +18,34 @@ def test_safe_replace_symbolic(): assert gtir_to_sdfg_utils.safe_replace_symbolic( dace.symbolic.pystr_to_symbolic("x*x + y"), symbol_mapping={"x": "y", "y": "x"} ) == dace.symbolic.pystr_to_symbolic("y*y + x") + + +def test_local_dimension_size(): + import numpy as np + + from gt4py._core import definitions as core_defs + from gt4py.next import common + from gt4py.next.program_processors.runners.dace import sdfg_args + + from next_tests.toy_connectivity import V2E, V2EDim, Vertex, Edge + + def conn_type(max_neighbors: int) -> common.NeighborConnectivityType: + return common.NeighborConnectivityType( + domain=(Vertex, V2EDim), + codomain=Edge, + skip_value=None, + dtype=core_defs.dtype(np.int32), + max_neighbors=max_neighbors, + ) + + sharer_tag = "some.module.V2EShared" + table_types = {V2E.offset_tag: conn_type(4), sharer_tag: conn_type(4)} + # a field finds the size in the table keyed by the local dimension + assert sdfg_args.local_dimension_size("a_field", V2EDim, table_types) == 4 + # a connectivity array has its own + conn_array = sdfg_args.connectivity_identifier(sharer_tag) + assert sdfg_args.local_dimension_size(conn_array, V2EDim, {sharer_tag: conn_type(4)}) == 4 + # a field over a local dimension bound only through a sharing connectivity + assert sdfg_args.local_dimension_size("a_field", V2EDim, {sharer_tag: conn_type(4)}) == 4 + with pytest.raises(KeyError, match="No connectivity over the local dimension"): + sdfg_args.local_dimension_size("a_field", V2EDim, {}) diff --git a/tests/next_tests/unit_tests/test_neighbor_connectivity.py b/tests/next_tests/unit_tests/test_neighbor_connectivity.py index 86bb41f1a3..a5d33b241c 100644 --- a/tests/next_tests/unit_tests/test_neighbor_connectivity.py +++ b/tests/next_tests/unit_tests/test_neighbor_connectivity.py @@ -146,7 +146,14 @@ class Local(DimensionIndex): ... class C(NeighborConnectivity[Vertex, Edge], max_neighbors=5): Local: typing.TypeAlias = V2E.Local """, - "counts are declared by its owner", + "contradicts the local dimension it shares with 'V2E'", + ), + ( + """ + class C(NeighborConnectivity[Edge, Edge]): + Local = V2E.Local + """, + "cannot share the local dimension of 'V2E'", ), ( """ @@ -249,6 +256,15 @@ class V2EShared(NeighborConnectivity[Vertex, Edge]): assert shared.offset_tag == shared.tag assert shared.__gt_type__().tag == shared.tag + def test_sharing_with_consistent_counts(self): + ns = _declare( + """ + class V2EShared4(NeighborConnectivity[Vertex, Edge], max_neighbors=4): + Local = V2E.Local + """ + ) + assert ns["V2EShared4"].Local is V2E.Local + def test_non_integer_index(self): with pytest.raises(TypeError, match="indexed by an integer"): V2E[Vertex] @@ -502,6 +518,46 @@ class Local(LocalDimensionIndex): ... assert not [w for w in recwarn if issubclass(w.category, DeprecationWarning)] +class TestConnectivityKeyOver: + def _type(self, connectivity): + return _table_type(domain=(connectivity.origin, common.local_dimension_of(connectivity))) + + def test_owner_is_preferred(self): + ns = _declare( + """ + class V2EShared(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + ) + shared = ns["V2EShared"] + provider = {shared.offset_tag: self._type(shared), V2E.offset_tag: self._type(V2E)} + assert common.connectivity_key_over(provider, V2E.Local) == V2E.offset_tag + assert common.connectivity_key_over(provider, V2E.Local.tag) == V2E.offset_tag + + def test_sharers_are_picked_independently_of_order(self): + ns = _declare( + """ + class SharedA(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + + class SharedB(NeighborConnectivity[Vertex, Edge]): + Local: typing.TypeAlias = V2E.Local + """ + ) + a, b = ns["SharedA"], ns["SharedB"] + forward = {a.offset_tag: self._type(a), b.offset_tag: self._type(b)} + backward = dict(reversed(forward.items())) + assert ( + common.connectivity_key_over(forward, V2E.Local) + == common.connectivity_key_over(backward, V2E.Local) + == min(a.offset_tag, b.offset_tag) + ) + + def test_nothing_bound(self): + with pytest.raises(KeyError, match="No connectivity over the local dimension"): + common.connectivity_key_over({E2V.offset_tag: self._type(E2V)}, V2E.Local) + + def test_the_const_list_dimension_cannot_be_adopted(): with pytest.raises(TypeError, match="cannot adopt"): _declare(