Skip to content
Merged
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
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -262,6 +262,8 @@ 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
4 changes: 3 additions & 1 deletion src/gt4py/next/ffront/fbuiltins.py
Original file line number Diff line number Diff line change
Expand Up @@ -480,6 +480,8 @@ def impl(
# guidelines for decision.
@dataclasses.dataclass(frozen=True)
class FieldOffset(runtime.Offset):
#: The tag, i.e. the offset-provider key.
value: str
source: common.Dimension
target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension]

Expand All @@ -492,7 +494,7 @@ def __post_init__(self) -> None:
raise ValueError("Second dimension in offset must be a local dimension.")

def __gt_type__(self) -> ts.OffsetType:
return ts.OffsetType(source=self.source, target=self.target)
return ts.OffsetType(source=self.source, target=self.target, tag=self.value)

def __getitem__(self, offset: int) -> common.Connectivity:
"""Serve as a connectivity factory."""
Expand Down
34 changes: 30 additions & 4 deletions src/gt4py/next/ffront/foast_passes/type_deduction.py
Original file line number Diff line number Diff line change
Expand Up @@ -456,13 +456,13 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri
f"Tuples need to be indexed with literal integers, got '{node.index}'.",
) from ex
new_type = types[index]
case ts.OffsetType(source=source, target=(target1, target2)):
case ts.OffsetType(source=source, target=(target1, target2), tag=tag):
if not target2.kind == DimensionKind.LOCAL:
raise errors.DSLError(
new_value.location, "Second dimension in offset must be a local dimension."
)
new_type = ts.OffsetType(source=source, target=(target1,))
case ts.OffsetType(source=source, target=(target,)):
new_type = ts.OffsetType(source=source, target=(target1,), tag=tag)
case ts.OffsetType(source=source, target=(target,), tag=tag):
# for cartesian axes (e.g. I, J) the index of the subscript only
# signifies the displacement in the respective dimension,
# but does not change the target type.
Expand All @@ -471,6 +471,19 @@ def visit_Subscript(self, node: foast.Subscript, **kwargs: Any) -> foast.Subscri
new_value.location,
"Source and target must be equal for offsets with a single target.",
)
if tag is None:
raise errors.DSLError(
new_value.location,
"Cannot index a dimension shift.",
notes=[
(
"A shift written as 'Dim + offset' already contains its"
" displacement, unlike a 'FieldOffset', which is indexed to"
" choose one."
)
],
hints=[f"Write the displacement directly, e.g. '{source.value} + 1'."],
)
new_type = new_value.type
case ts.FieldType(dims=dims, dtype=dtype):
# e.g. `field[LocalDim(42)]`
Expand Down Expand Up @@ -769,7 +782,20 @@ def visit_Call(self, node: foast.Call, **kwargs: Any) -> foast.Call:
):
raise errors.DSLError(node.location, "Functions can only be called directly.")
elif isinstance(new_func.type, ts.FieldType):
pass
for arg in new_args:
# A Cartesian `FieldOffset` shifts by the index it is subscripted with, so it
# carries no displacement on its own. Only an offset with a local dimension is
# meaningful unsubscripted, as the neighbor access `field(Off)`.
if (
isinstance(arg, (foast.Name, foast.Attribute))
and isinstance(arg.type, ts.OffsetType)
and len(arg.type.target) == 1
):
raise errors.DSLError(
arg.location,
f"Cannot shift by the Cartesian offset '{arg!s}' without an index.",
hints=[f"Give the displacement, e.g. '{arg!s}[1]'."],
)
elif isinstance(new_func.type, ts.DimensionType):
assert new_func.type.dim.kind == DimensionKind.LOCAL
return foast.Call(
Expand Down
23 changes: 12 additions & 11 deletions src/gt4py/next/ffront/foast_to_gtir.py
Original file line number Diff line number Diff line change
Expand Up @@ -295,14 +295,17 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr:
for arg in node.args:
match arg:
# `field(Off[idx])`
case foast.Subscript(value=foast.Name() as offset_name, index=index):
# (matched on the type, not the node, to also accept `mod.Off[idx]`)
case foast.Subscript(
value=foast.LocatedNode(type=ts.OffsetType(tag=str() as offset_tag)),
index=index,
):
# Constant folding to a `Literal` ensures that `index` becomes an `OffsetLiteral`,
# which can be generated as compile-time value backend code.
new_index = constant_folding.ConstantFolding.apply(self.visit(index, **kwargs))
assert isinstance(new_index, itir.Literal)
assert isinstance(offset_name.type, ts.OffsetType)
current_expr = im.as_fieldop(
im.lambda_("__it")(im.deref(im.shift(offset_name.id, new_index)("__it")))
im.lambda_("__it")(im.deref(im.shift(offset_tag, new_index)("__it")))
)(current_expr)
# `field(Dim + idx)` (where `idx` is integer or half integer)
case foast.BinOp(
Expand All @@ -322,14 +325,6 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr:
)
)
)(current_expr)
# `field(Off)`
case foast.Name(id=offset_name):
# only a single unstructured shift is supported so returning here is fine even though we
# are in a loop.
assert len(node.args) == 1 and len(arg.type.target) > 1 # type: ignore[attr-defined] # ensured by pattern
return im.as_fieldop_neighbors(
str(offset_name), self.visit(node.func, **kwargs)
)
# `field(as_offset(Off, offset_field))`
case foast.Call(func=foast.Name(id="as_offset")):
func_args = arg
Expand All @@ -344,6 +339,12 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr:
)
)
)(current_expr, offset_field)
# `field(Off)`
case foast.LocatedNode(type=ts.OffsetType(tag=str() as offset_tag, target=(_, _))):
# only a single unstructured shift is supported so returning here is fine even though we
# are in a loop.
assert len(node.args) == 1
return im.as_fieldop_neighbors(offset_tag, self.visit(node.func, **kwargs))
case _:
raise FieldOperatorLoweringError("Unexpected shift arguments!")

Expand Down
5 changes: 4 additions & 1 deletion src/gt4py/next/type_system/type_specifications.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,9 +74,12 @@ class OffsetType(TypeSpec):
# TODO(havogt): replace by ConnectivityType
source: common.Dimension
target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension]
#: The offset-provider key; `None` for the untagged Cartesian `Dim + offset`.
tag: Optional[common.Tag] = None

def __str__(self) -> str:
return f"Offset[{self.source}, {self.target}]"
tag = "" if self.tag is None else f"{self.tag}: "
return f"Offset[{tag}{self.source}, {self.target}]"


class ScalarKind(eve_types.IntEnum):
Expand Down
43 changes: 42 additions & 1 deletion tests/next_tests/definitions.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,10 @@ 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"
Expand Down Expand Up @@ -134,6 +138,14 @@ class ProgramFormatterId(_PythonObjectIdMixin, str, enum.Enum):
REDUCTION_WITH_ONLY_SPARSE_FIELDS_MESSAGE = (
"We cannot unroll a reduction on a sparse field only (not clear if it is legal ITIR)"
)
#: An offset and its local dimension must currently share a name on most backends, because
#: 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
Expand Down Expand Up @@ -169,6 +181,16 @@ 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 = [
Expand All @@ -180,6 +202,11 @@ 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),
Expand All @@ -190,7 +217,15 @@ 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 + []
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,
),
]
GTFN_SKIP_TEST_LIST = (
COMMON_SKIP_TEST_LIST
+ DOMAIN_INFERENCE_SKIP_LIST
Expand All @@ -201,6 +236,12 @@ 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,
),
]
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,11 @@
import numpy as np

import gt4py.next as gtx
from gt4py.next import broadcast, astype, int32
from gt4py.next import broadcast, astype, common, int32, neighbor_sum

from next_tests import integration_tests
from next_tests.integration_tests import cases
from next_tests.integration_tests.cases import cartesian_case, IDim, KDim
from next_tests.integration_tests.cases import cartesian_case, unstructured_case, IDim, KDim

from next_tests.integration_tests.cases_utils import (
exec_alloc_descriptor,
Expand Down Expand Up @@ -54,6 +54,33 @@ def mod_prog(f: cases.IField, isize: int32, ksize: int32, out: cases.IKField):
cases.verify(cartesian_case, mod_prog, f, isize, ksize, out=out, ref=expected)


@pytest.mark.uses_unstructured_shift
def test_import_offset_module_unstructured_shift(unstructured_case):
@gtx.field_operator
def testee(a: cases.EField) -> cases.VField:
return neighbor_sum(a(cases.V2E), axis=cases.V2EDim)

v2e_table = unstructured_case.offset_provider["V2E"].asnumpy()
cases.verify_with_default_data(
unstructured_case,
testee,
ref=lambda a: np.sum(a[v2e_table], axis=1, where=v2e_table != common._DEFAULT_SKIP_VALUE),
)


@pytest.mark.uses_unstructured_shift
def test_import_offset_module_sparse_shift(unstructured_case):
@gtx.field_operator
def testee(a: cases.VField) -> cases.EField:
return a(cases.E2V[0])

cases.verify_with_default_data(
unstructured_case,
testee,
ref=lambda a: a[unstructured_case.offset_provider["E2V"].asnumpy()[:, 0]],
)


# TODO: these set of features should be allowed as module imports in a later PR
def test_import_module_errors_future_allowed(cartesian_case):
from ....artifacts.dummy_package import dummy_module
Expand Down
Loading
Loading