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/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..5ffc813f8c 100644 --- a/src/gt4py/next/common.py +++ b/src/gt4py/next/common.py @@ -1471,6 +1471,48 @@ 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 | Tag +) -> str: + """ + 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 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. + """ + 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: """Determine if offset provider has an element for the given offset tag.""" try: @@ -2128,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/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..7062053223 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, tag), + ) assert common.is_neighbor_table(offset_implementation) source_dim = offset_implementation.__gt_type__().source_dim cur_index = pos[source_dim.tag] @@ -1007,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.offset.value: i}) # type: ignore[dict-item] + self._translate_named_indices({**named_indices, local_tag: i}) ] = v elif isinstance(value, _ConstList): self._ndarrayfield[ @@ -1420,16 +1426,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) @@ -1452,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 @@ -1485,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) @@ -1543,7 +1561,10 @@ 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, self.list_offset) + connectivity = common.get_offset(offset_provider, connectivity_key) assert common.is_neighbor_table(connectivity) return _List( values=tuple( @@ -1553,7 +1574,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 +1821,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/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/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_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 924758c727..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 @@ -1321,7 +1323,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 +1379,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 @@ -1437,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 @@ -1477,7 +1483,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 +1497,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/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/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..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 @@ -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,57 @@ 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), + ) + + +@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/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): 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(