diff --git a/src/gt4py/next/ffront/fbuiltins.py b/src/gt4py/next/ffront/fbuiltins.py index 602221fbcc..955554dc43 100644 --- a/src/gt4py/next/ffront/fbuiltins.py +++ b/src/gt4py/next/ffront/fbuiltins.py @@ -482,7 +482,8 @@ 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) + assert isinstance(self.value, str) + return ts.OffsetType(source=self.source, target=self.target, name=self.value) def __getitem__(self, offset: int) -> common.Connectivity: """Serve as a connectivity factory.""" diff --git a/src/gt4py/next/ffront/foast_passes/type_deduction.py b/src/gt4py/next/ffront/foast_passes/type_deduction.py index c9f51ad080..f824501e2b 100644 --- a/src/gt4py/next/ffront/foast_passes/type_deduction.py +++ b/src/gt4py/next/ffront/foast_passes/type_deduction.py @@ -456,12 +456,12 @@ 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), name=name): 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,)) + new_type = ts.OffsetType(source=source, target=(target1,), name=name) case ts.OffsetType(source=source, target=(target,)): # for cartesian axes (e.g. I, J) the index of the subscript only # signifies the displacement in the respective dimension, diff --git a/src/gt4py/next/ffront/foast_to_gtir.py b/src/gt4py/next/ffront/foast_to_gtir.py index 480a05812f..eda4c2e6b1 100644 --- a/src/gt4py/next/ffront/foast_to_gtir.py +++ b/src/gt4py/next/ffront/foast_to_gtir.py @@ -295,14 +295,16 @@ 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): + case foast.Subscript( + value=foast.LocatedNode(type=ts.OffsetType(name=str() as offset_name)), + 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_name, new_index)("__it"))) )(current_expr) # `field(Dim + idx)` (where `idx` is integer or half integer) case foast.BinOp( @@ -322,14 +324,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 @@ -344,6 +338,14 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr: ) ) )(current_expr, offset_field) + # `field(Off)` + case foast.LocatedNode( + type=ts.OffsetType(name=str() as offset_name, 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_name, self.visit(node.func, **kwargs)) case _: raise FieldOperatorLoweringError("Unexpected shift arguments!") diff --git a/src/gt4py/next/type_system/type_specifications.py b/src/gt4py/next/type_system/type_specifications.py index 59ac40f0f3..4afe697f0d 100644 --- a/src/gt4py/next/type_system/type_specifications.py +++ b/src/gt4py/next/type_system/type_specifications.py @@ -74,6 +74,9 @@ class OffsetType(TypeSpec): # TODO(havogt): replace by ConnectivityType source: common.Dimension target: tuple[common.Dimension] | tuple[common.Dimension, common.Dimension] + #: `None` for the offsets synthesized from `Dim + idx` dimension arithmetic, which are + #: resolved from `source`/`target` alone rather than looked up in the offset provider. + name: Optional[common.Tag] = None def __str__(self) -> str: return f"Offset[{self.source}, {self.target}]" diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py index 899197c43f..67dcf93f8e 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_import_from_mod.py @@ -10,11 +10,19 @@ 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, + IHalfDim, + KDim, + V2EDim, +) +from next_tests.integration_tests.cases import E2V as RenamedE2V, V2E as RenamedV2E from next_tests.integration_tests.cases_utils import ( exec_alloc_descriptor, @@ -54,6 +62,73 @@ 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_cartesian_shift +def test_import_dims_module_cartesian_shift(cartesian_case): + @gtx.field_operator + def testee(a: cases.IField) -> cases.IField: + return a(cases.IDim + 1) + + size = cartesian_case.default_sizes[IDim] + a = cases.allocate(cartesian_case, testee, "a", domain={IDim: (1, size + 1)})() + out = cases.allocate(cartesian_case, testee, cases.RETURN, domain={IDim: (0, size)})() + + cases.verify(cartesian_case, testee, a, out=out, ref=a[:], offset_provider={}) + + +@pytest.mark.uses_cartesian_shift +def test_import_dims_module_staggered_shift(cartesian_case): + @gtx.field_operator + def testee(a: cases.IField) -> cases.IHalfField: + return a(cases.IHalfDim + 0.5) + + size = cartesian_case.default_sizes[IDim] + a = cases.allocate(cartesian_case, testee, "a", sizes={IDim: size})() + out = cases.allocate(cartesian_case, testee, cases.RETURN, sizes={IHalfDim: size})() + + cases.verify(cartesian_case, testee, a, out=out, ref=a, offset_provider={}) + + +@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_renamed_offset_unstructured_shift(unstructured_case): + @gtx.field_operator + def testee(a: cases.EField) -> cases.VField: + return neighbor_sum(a(RenamedV2E), axis=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_renamed_offset_sparse_shift(unstructured_case): + @gtx.field_operator + def testee(a: cases.VField) -> cases.EField: + return a(RenamedE2V[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