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
4 changes: 2 additions & 2 deletions src/gt4py/next/ffront/foast_to_gtir.py
Original file line number Diff line number Diff line change
Expand Up @@ -322,8 +322,8 @@ def _visit_shift(self, node: foast.Call, **kwargs: Any) -> itir.Expr:
)
)
)(current_expr)
# `field(Off)`
case foast.Name(id=offset_name):
# `field(Off)` or `field(mod.Off)`
case foast.Name(id=offset_name) | foast.Attribute(attr=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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,17 @@
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,
)

from next_tests.integration_tests.cases_utils import (
exec_alloc_descriptor,
Expand Down Expand Up @@ -54,6 +60,46 @@ 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),
)


# 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