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
3 changes: 2 additions & 1 deletion src/gt4py/next/ffront/fbuiltins.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
4 changes: 2 additions & 2 deletions src/gt4py/next/ffront/foast_passes/type_deduction.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
24 changes: 13 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,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(
Expand All @@ -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
Expand All @@ -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!")

Expand Down
3 changes: 3 additions & 0 deletions src/gt4py/next/type_system/type_specifications.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}]"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down