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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 14 additions & 9 deletions docs/development/ADRs/next/0029-Connectivities_As_Types.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 0 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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',
Expand Down
63 changes: 57 additions & 6 deletions src/gt4py/next/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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 (
Expand Down
4 changes: 2 additions & 2 deletions src/gt4py/next/embedded/nd_array_field.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])

Expand Down
53 changes: 38 additions & 15 deletions src/gt4py/next/iterator/embedded.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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[
Expand Down Expand Up @@ -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)
Expand All @@ -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


Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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(
Expand All @@ -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:
Expand Down Expand Up @@ -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),
Expand Down
2 changes: 1 addition & 1 deletion src/gt4py/next/iterator/tracing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__ == "<lambda>":
return lambdadef(o)
Expand Down
10 changes: 6 additions & 4 deletions src/gt4py/next/iterator/transforms/unroll_reduce.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
]
Expand All @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -91,7 +91,12 @@ def _process_regular_arguments(
# NOTE: the tag is the offset-provider key, and its mangled form names the
# `generated::<name>_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<generated::{common.codegen_name(dim_name)}_t, {size}>({arg})"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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: ...

Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading
Loading