From bb4bfdc9727b9395fd121ec5d5fe301cb6da54a3 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 10:15:38 +0200 Subject: [PATCH 1/6] build: pin gt4py to the head of the connectivities-as-types stack Temporary [tool.uv.sources] override to GridTools/gt4py#2912 at 22aab76b9, the last PR of stack #2917 (#2899 -> #2907 -> #2910 -> #2912). Pinned by rev, not branch, so a rebase of the stack cannot silently move what this is built against. Revert once the stack is released. The manifest `gt4py==` pins are left alone; a source override does not enforce them. gt4py pins `dace==2.0.0a9`, which moves dace from 2.0.0a7; nothing else in the resolution changes. --- pyproject.toml | 2 +- uv.lock | 45 +++++++++++++++++++++------------------------ 2 files changed, 22 insertions(+), 25 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 96f884a238..48d56e096d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -423,7 +423,7 @@ url = 'https://gridtools.github.io/pypi/' [tool.uv.sources] # dace = {index = "gridtools"} -# gt4py = {git = "https://github.com/GridTools/gt4py", branch = "main"} +gt4py = {git = "https://github.com/GridTools/gt4py", rev = "22aab76b917535670b58a1d6774f6cbddaa2299b"} # gt4py = {index = "test.pypi"} icon4py-atmosphere-diffusion = {workspace = true} icon4py-atmosphere-dycore = {workspace = true} diff --git a/uv.lock b/uv.lock index 3f3ae2c716..61215ab54a 100644 --- a/uv.lock +++ b/uv.lock @@ -1163,6 +1163,7 @@ wheels = [ [[package]] dependencies = [ {name = "astunparse"}, + {name = "click"}, {name = "cmake"}, {name = "dill"}, {name = "fparser"}, @@ -1180,9 +1181,9 @@ dependencies = [ {name = "typing-extensions"} ] name = "dace" -sdist = {url = "https://files.pythonhosted.org/packages/df/e0/2c45877476eb300fc8211fb2a46b3e3203468cd2bd3e0ecb4a7f5a7a4694/dace-2.0.0a7.tar.gz", hash = "sha256:ab15a73ce4c2cacaf9244832f84ea56328e181b21956a6d0c529782201609095", size = 6827344, upload-time = "2026-08-28T14:11:38.633Z"} +sdist = {url = "https://files.pythonhosted.org/packages/76/34/a6e66f329ea94167c1c0da89dfaf053b1d90836ccab7161176dd12f10f9e/dace-2.0.0a9.tar.gz", hash = "sha256:ee9cd2a40c16fb09d559556433bb4f51b71fa88760fbc04306d54552775b5c4a", size = 6572070, upload-time = "2026-09-17T09:31:23.155Z"} source = {registry = "https://pypi.org/simple"} -version = "2.0.0a7" +version = "2.0.0a9" [[package]] dependencies = [ @@ -1660,12 +1661,8 @@ dependencies = [ {name = "xxhash"} ] name = "gt4py" -sdist = {url = "https://files.pythonhosted.org/packages/b4/54/c9a0cfe5eb6289b6cc2a63113e9bff4807dfc672b2512f0bac5ac74d8740/gt4py-1.2.2.tar.gz", hash = "sha256:78704e71fb719d7d5ba61690f9ea749e80c5088ef194ffdb3d5c4bff31585739", size = 891724, upload-time = "2026-09-02T08:10:11.521Z"} -source = {registry = "https://pypi.org/simple"} -version = "1.2.2" -wheels = [ - {url = "https://files.pythonhosted.org/packages/30/dc/a66802bed6b499349315b5e20108341306d0430235e1b1d0e56cf011e97e/gt4py-1.2.2-py3-none-any.whl", hash = "sha256:2aa7ea57081ed1ba51b202bc87e89a67c5fcb3588a8cb9b3a13d5343f0d83762", size = 1111109, upload-time = "2026-09-02T08:10:09.771Z"} -] +source = {git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b#22aab76b917535670b58a1d6774f6cbddaa2299b"} +version = "1.2.2.post56+22aab76b" [package.optional-dependencies] cuda12 = [ @@ -2050,7 +2047,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", editable = "model/common"}, {name = "packaging", specifier = ">=20.0"} ] @@ -2067,7 +2064,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", editable = "model/common"}, {name = "packaging", specifier = ">=20.0"} ] @@ -2084,7 +2081,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", editable = "model/common"}, {name = "packaging", specifier = ">=20.0"} ] @@ -2102,7 +2099,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", extras = ["io"], editable = "model/common"}, {name = "numpy", specifier = ">=1.23.3"}, {name = "packaging", specifier = ">=20.0"} @@ -2135,7 +2132,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", editable = "model/common"}, {name = "packaging", specifier = ">=20.0"} ] @@ -2152,7 +2149,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", editable = "model/common"}, {name = "packaging", specifier = ">=20.0"} ] @@ -2180,9 +2177,9 @@ requires-dist = [ {name = "click", specifier = ">=8.0"}, {name = "cupy-cuda12x", marker = "extra == 'cuda12'", specifier = ">=13.0"}, {name = "cupy-cuda13x", marker = "extra == 'cuda13'", specifier = ">=14.0"}, - {name = "gt4py", specifier = "==1.2.2"}, - {name = "gt4py", extras = ["cuda12"], marker = "extra == 'cuda12'"}, - {name = "gt4py", extras = ["cuda13"], marker = "extra == 'cuda13'"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, + {name = "gt4py", extras = ["cuda12"], marker = "extra == 'cuda12'", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, + {name = "gt4py", extras = ["cuda13"], marker = "extra == 'cuda13'", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-atmosphere-diffusion", editable = "model/atmosphere/diffusion"}, {name = "icon4py-atmosphere-dycore", editable = "model/atmosphere/dycore"}, {name = "icon4py-common", extras = ["distributed"], editable = "model/common"}, @@ -2231,10 +2228,10 @@ requires-dist = [ {name = "cupy-rocm-7-0", marker = "extra == 'rocm7'", specifier = ">=14.1.1"}, {name = "datashader", marker = "extra == 'io'", specifier = ">=0.16.1"}, {name = "ghex", marker = "extra == 'distributed'", specifier = ">=0.9.0"}, - {name = "gt4py", specifier = "==1.2.2"}, - {name = "gt4py", extras = ["cuda12"], marker = "extra == 'cuda12'"}, - {name = "gt4py", extras = ["cuda13"], marker = "extra == 'cuda13'"}, - {name = "gt4py", extras = ["rocm7"], marker = "extra == 'rocm7'"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, + {name = "gt4py", extras = ["cuda12"], marker = "extra == 'cuda12'", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, + {name = "gt4py", extras = ["cuda13"], marker = "extra == 'cuda13'", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, + {name = "gt4py", extras = ["rocm7"], marker = "extra == 'rocm7'", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "holoviews", marker = "extra == 'io'", specifier = ">=1.16.0"}, {name = "icon4py-common", extras = ["distributed", "io"], marker = "extra == 'all'", editable = "model/common"}, {name = "mpi4py", marker = "extra == 'distributed'", specifier = ">=3.1.5"}, @@ -2318,7 +2315,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ {name = "devtools", specifier = ">=0.12"}, - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-atmosphere-diffusion", editable = "model/atmosphere/diffusion"}, {name = "icon4py-atmosphere-dycore", editable = "model/atmosphere/dycore"}, {name = "icon4py-atmosphere-microphysics", editable = "model/atmosphere/subgrid_scale_physics/microphysics"}, @@ -2349,7 +2346,7 @@ version = "0.4.1" [package.metadata] requires-dist = [ {name = "filelock", specifier = ">=3.18.0,<3.20"}, - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", extras = ["io"], editable = "model/common"}, {name = "icon4py-driver", editable = "model/driver"}, {name = "numpy", specifier = ">=1.23.3"}, @@ -2382,7 +2379,7 @@ requires-dist = [ {name = "click", specifier = ">=8.1.7"}, {name = "configargparse", specifier = ">=1.7.1"}, {name = "fprettify", specifier = ">=0.3.7"}, - {name = "gt4py", specifier = "==1.2.2"}, + {name = "gt4py", git = "https://github.com/GridTools/gt4py?rev=22aab76b917535670b58a1d6774f6cbddaa2299b"}, {name = "icon4py-common", extras = ["cuda12"], marker = "extra == 'cuda12'", editable = "model/common"}, {name = "icon4py-common", extras = ["cuda13"], marker = "extra == 'cuda13'", editable = "model/common"}, {name = "icon4py-common", extras = ["rocm7"], marker = "extra == 'rocm7'", editable = "model/common"}, From adbf9d7383ca92c508e483d803636862dfa0c7cb Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 10:15:48 +0200 Subject: [PATCH 2/6] refactor: declare dimensions and connectivities as classes Output of gt4py's `scripts/python/migrate_connectivities.py` (from #2910, at the pinned rev) run over model/, tools/ and bindings/, then `ruff format`. Dimensions become `DimensionIndex` / `LocalDimensionIndex` subclasses, whose tag is now their qualified name. Each `FieldOffset` becomes a `NeighborConnectivity[Domain, Codomain]` that adopts the existing local dimension (`Local: typing.TypeAlias = E2CDim`), so every `dims.XxxDim` stays valid. `Koff` is removed and its uses become `KDim + i` / `as_offset(KDim, ...)`. `Local` must stay an annotated `TypeAlias`: the PEP 695 form ruff's UP040 asks for yields a `TypeAliasType` rather than the class, so the rule is silenced on those lines. --- .../tests/bindings/test_icon4py_export.py | 3 +- ...fusion_nabla_of_theta_over_steep_points.py | 14 +- ...t_of_exner_pressure_for_multiple_levels.py | 14 +- .../compute_hydrostatic_correction_term.py | 18 +-- .../src/icon4py/model/common/dimension.py | 148 ++++++++++++++---- 5 files changed, 139 insertions(+), 58 deletions(-) diff --git a/bindings/tests/bindings/test_icon4py_export.py b/bindings/tests/bindings/test_icon4py_export.py index b7791edee2..f781a37e47 100644 --- a/bindings/tests/bindings/test_icon4py_export.py +++ b/bindings/tests/bindings/test_icon4py_export.py @@ -33,7 +33,8 @@ }, ) -SomeDim = gtx.Dimension("SomeDim") + +class SomeDim(gtx.DimensionIndex): ... def make_array_info( diff --git a/model/atmosphere/diffusion/src/icon4py/model/atmosphere/diffusion/stencils/truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py b/model/atmosphere/diffusion/src/icon4py/model/atmosphere/diffusion/stencils/truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py index 04974649d3..7a1b679f57 100644 --- a/model/atmosphere/diffusion/src/icon4py/model/atmosphere/diffusion/stencils/truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py +++ b/model/atmosphere/diffusion/src/icon4py/model/atmosphere/diffusion/stencils/truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py @@ -10,7 +10,7 @@ from gt4py.next.experimental import as_offset from icon4py.model.common import dimension as dims, field_type_aliases as fa -from icon4py.model.common.dimension import C2E2C, Koff +from icon4py.model.common.dimension import C2E2C, KDim from icon4py.model.common.type_alias import vpfloat, wpfloat @@ -26,21 +26,21 @@ def _truly_horizontal_diffusion_nabla_of_theta_over_steep_points( ) -> fa.CellKField[vpfloat]: z_temp_wp = astype(z_temp, wpfloat) - theta_v_0 = theta_v(C2E2C[0])(as_offset(Koff, zd_vertoffset[dims.C2E2CDim(0)])) - theta_v_1 = theta_v(C2E2C[1])(as_offset(Koff, zd_vertoffset[dims.C2E2CDim(1)])) - theta_v_2 = theta_v(C2E2C[2])(as_offset(Koff, zd_vertoffset[dims.C2E2CDim(2)])) + theta_v_0 = theta_v(C2E2C[0])(as_offset(KDim, zd_vertoffset[dims.C2E2CDim(0)])) + theta_v_1 = theta_v(C2E2C[1])(as_offset(KDim, zd_vertoffset[dims.C2E2CDim(1)])) + theta_v_2 = theta_v(C2E2C[2])(as_offset(KDim, zd_vertoffset[dims.C2E2CDim(2)])) # `zd_vertoffset` is 0 where `zd_diffcoef` is 0, so there the `+ 1` target at the bottom level # lies below the column; keep the read in range for backends that evaluate every point. is_steep = zd_diffcoef != wpfloat("0.0") theta_v_0_m1 = theta_v(C2E2C[0])( - as_offset(Koff, where(is_steep, zd_vertoffset[dims.C2E2CDim(0)] + 1, 0)) + as_offset(KDim, where(is_steep, zd_vertoffset[dims.C2E2CDim(0)] + 1, 0)) ) theta_v_1_m1 = theta_v(C2E2C[1])( - as_offset(Koff, where(is_steep, zd_vertoffset[dims.C2E2CDim(1)] + 1, 0)) + as_offset(KDim, where(is_steep, zd_vertoffset[dims.C2E2CDim(1)] + 1, 0)) ) theta_v_2_m1 = theta_v(C2E2C[2])( - as_offset(Koff, where(is_steep, zd_vertoffset[dims.C2E2CDim(2)] + 1, 0)) + as_offset(KDim, where(is_steep, zd_vertoffset[dims.C2E2CDim(2)] + 1, 0)) ) sum_tmp = ( diff --git a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py index cce0eb7d5d..af8522e101 100644 --- a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py +++ b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py @@ -10,7 +10,7 @@ from gt4py.next.experimental import as_offset from icon4py.model.common import dimension as dims, field_type_aliases as fa -from icon4py.model.common.dimension import E2C, Koff +from icon4py.model.common.dimension import E2C, KDim from icon4py.model.common.type_alias import vpfloat, wpfloat @@ -24,14 +24,14 @@ def _compute_horizontal_gradient_of_exner_pressure_for_multiple_levels( z_dexner_dz_c_2: fa.CellKField[vpfloat], ) -> fa.EdgeKField[vpfloat]: """Formerly known as _mo_solve_nonhydro_stencil_20.""" - z_exner_ex_pr_0 = z_exner_ex_pr(E2C[0])(as_offset(Koff, ikoffset[dims.E2CDim(0)])) - z_exner_ex_pr_1 = z_exner_ex_pr(E2C[1])(as_offset(Koff, ikoffset[dims.E2CDim(1)])) + z_exner_ex_pr_0 = z_exner_ex_pr(E2C[0])(as_offset(KDim, ikoffset[dims.E2CDim(0)])) + z_exner_ex_pr_1 = z_exner_ex_pr(E2C[1])(as_offset(KDim, ikoffset[dims.E2CDim(1)])) - z_dexner_dz_c1_0 = z_dexner_dz_c_1(E2C[0])(as_offset(Koff, ikoffset[dims.E2CDim(0)])) - z_dexner_dz_c1_1 = z_dexner_dz_c_1(E2C[1])(as_offset(Koff, ikoffset[dims.E2CDim(1)])) + z_dexner_dz_c1_0 = z_dexner_dz_c_1(E2C[0])(as_offset(KDim, ikoffset[dims.E2CDim(0)])) + z_dexner_dz_c1_1 = z_dexner_dz_c_1(E2C[1])(as_offset(KDim, ikoffset[dims.E2CDim(1)])) - z_dexner_dz_c2_0 = z_dexner_dz_c_2(E2C[0])(as_offset(Koff, ikoffset[dims.E2CDim(0)])) - z_dexner_dz_c2_1 = z_dexner_dz_c_2(E2C[1])(as_offset(Koff, ikoffset[dims.E2CDim(1)])) + z_dexner_dz_c2_0 = z_dexner_dz_c_2(E2C[0])(as_offset(KDim, ikoffset[dims.E2CDim(0)])) + z_dexner_dz_c2_1 = z_dexner_dz_c_2(E2C[1])(as_offset(KDim, ikoffset[dims.E2CDim(1)])) z_gradh_exner_wp = inv_dual_edge_length * ( astype( diff --git a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_hydrostatic_correction_term.py b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_hydrostatic_correction_term.py index 34e3bd5873..cab4bf50b9 100644 --- a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_hydrostatic_correction_term.py +++ b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/compute_hydrostatic_correction_term.py @@ -10,7 +10,7 @@ from gt4py.next.experimental import as_offset from icon4py.model.common import dimension as dims, field_type_aliases as fa -from icon4py.model.common.dimension import E2C, Koff +from icon4py.model.common.dimension import E2C, KDim from icon4py.model.common.type_alias import vpfloat, wpfloat @@ -27,27 +27,27 @@ def _compute_hydrostatic_correction_term( """Formerly known as _mo_solve_nonhydro_stencil_21.""" zdiff_gradp_wp = zdiff_gradp # astype(zdiff_gradp, wpfloat) # TODO(): fix this cast - theta_v_0 = theta_v(E2C[0])(as_offset(Koff, ikoffset[dims.E2CDim(0)])) - theta_v_1 = theta_v(E2C[1])(as_offset(Koff, ikoffset[dims.E2CDim(1)])) + theta_v_0 = theta_v(E2C[0])(as_offset(KDim, ikoffset[dims.E2CDim(0)])) + theta_v_1 = theta_v(E2C[1])(as_offset(KDim, ikoffset[dims.E2CDim(1)])) # Each `theta_v_ic` chain must stay written out. Binding the shifted access to a local # emits a shared `let` that the dace lowering of `as_offset` cannot fuse, and the program # then fails to compile. - theta_v_ic_0 = theta_v_ic(E2C[0])(dims.KDim - 0.5)(as_offset(Koff, ikoffset[dims.E2CDim(0)])) - theta_v_ic_1 = theta_v_ic(E2C[1])(dims.KDim - 0.5)(as_offset(Koff, ikoffset[dims.E2CDim(1)])) + theta_v_ic_0 = theta_v_ic(E2C[0])(dims.KDim - 0.5)(as_offset(KDim, ikoffset[dims.E2CDim(0)])) + theta_v_ic_1 = theta_v_ic(E2C[1])(dims.KDim - 0.5)(as_offset(KDim, ikoffset[dims.E2CDim(1)])) theta_v_ic_p1_0 = theta_v_ic(E2C[0])(dims.KDim - 0.5)( - as_offset(Koff, ikoffset[dims.E2CDim(0)] + 1) + as_offset(KDim, ikoffset[dims.E2CDim(0)] + 1) ) theta_v_ic_p1_1 = theta_v_ic(E2C[1])(dims.KDim - 0.5)( - as_offset(Koff, ikoffset[dims.E2CDim(1)] + 1) + as_offset(KDim, ikoffset[dims.E2CDim(1)] + 1) ) inv_ddqz_z_full_0_wp = astype( - inv_ddqz_z_full(E2C[0])(as_offset(Koff, ikoffset[dims.E2CDim(0)])), wpfloat + inv_ddqz_z_full(E2C[0])(as_offset(KDim, ikoffset[dims.E2CDim(0)])), wpfloat ) inv_ddqz_z_full_1_wp = astype( - inv_ddqz_z_full(E2C[1])(as_offset(Koff, ikoffset[dims.E2CDim(1)])), wpfloat + inv_ddqz_z_full(E2C[1])(as_offset(KDim, ikoffset[dims.E2CDim(1)])), wpfloat ) z_theta_0 = ( diff --git a/model/common/src/icon4py/model/common/dimension.py b/model/common/src/icon4py/model/common/dimension.py index f4d47564d0..3fdead524e 100644 --- a/model/common/src/icon4py/model/common/dimension.py +++ b/model/common/src/icon4py/model/common/dimension.py @@ -5,46 +5,126 @@ # # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import typing from collections.abc import Iterator import gt4py.next as gtx -KDim = gtx.Dimension("K", kind=gtx.DimensionKind.VERTICAL) +class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + KHalfDim = gtx.flip_staggered(KDim) -EdgeDim = gtx.Dimension("Edge") -CellDim = gtx.Dimension("Cell") -VertexDim = gtx.Dimension("Vertex") -LsqUnkDim = gtx.Dimension("LsqUnk", gtx.DimensionKind.LOCAL) -E2CDim = gtx.Dimension("E2C", gtx.DimensionKind.LOCAL) -E2VDim = gtx.Dimension("E2V", gtx.DimensionKind.LOCAL) -C2EDim = gtx.Dimension("C2E", gtx.DimensionKind.LOCAL) -V2CDim = gtx.Dimension("V2C", gtx.DimensionKind.LOCAL) -C2VDim = gtx.Dimension("C2V", gtx.DimensionKind.LOCAL) -V2EDim = gtx.Dimension("V2E", gtx.DimensionKind.LOCAL) -V2E2VDim = gtx.Dimension("V2E2V", gtx.DimensionKind.LOCAL) -E2C2VDim = gtx.Dimension("E2C2V", gtx.DimensionKind.LOCAL) -C2E2CODim = gtx.Dimension("C2E2CO", gtx.DimensionKind.LOCAL) -E2C2EODim = gtx.Dimension("E2C2EO", gtx.DimensionKind.LOCAL) -E2C2EDim = gtx.Dimension("E2C2E", gtx.DimensionKind.LOCAL) -C2E2CDim = gtx.Dimension("C2E2C", gtx.DimensionKind.LOCAL) -C2E2C2EDim = gtx.Dimension("C2E2C2E", gtx.DimensionKind.LOCAL) -C2E2C2E2CDim = gtx.Dimension("C2E2C2E2C", gtx.DimensionKind.LOCAL) -E2C = gtx.FieldOffset("E2C", source=CellDim, target=(EdgeDim, E2CDim)) -C2E = gtx.FieldOffset("C2E", source=EdgeDim, target=(CellDim, C2EDim)) -V2C = gtx.FieldOffset("V2C", source=CellDim, target=(VertexDim, V2CDim)) -C2V = gtx.FieldOffset("C2V", source=VertexDim, target=(CellDim, C2VDim)) -V2E = gtx.FieldOffset("V2E", source=EdgeDim, target=(VertexDim, V2EDim)) -E2V = gtx.FieldOffset("E2V", source=VertexDim, target=(EdgeDim, E2VDim)) -E2C2V = gtx.FieldOffset("E2C2V", source=VertexDim, target=(EdgeDim, E2C2VDim)) -C2E2CO = gtx.FieldOffset("C2E2CO", source=CellDim, target=(CellDim, C2E2CODim)) -E2C2EO = gtx.FieldOffset("E2C2EO", source=EdgeDim, target=(EdgeDim, E2C2EODim)) -E2C2E = gtx.FieldOffset("E2C2E", source=EdgeDim, target=(EdgeDim, E2C2EDim)) -C2E2C = gtx.FieldOffset("C2E2C", source=CellDim, target=(CellDim, C2E2CDim)) -C2E2C2E = gtx.FieldOffset("C2E2C2E", source=EdgeDim, target=(CellDim, C2E2C2EDim)) -C2E2C2E2C = gtx.FieldOffset("C2E2C2E2C", source=CellDim, target=(CellDim, C2E2C2E2CDim)) -V2E2V = gtx.FieldOffset("V2E2V", source=VertexDim, target=(VertexDim, V2E2VDim)) -Koff = gtx.FieldOffset("Koff", source=KDim, target=(KDim,)) + + +class EdgeDim(gtx.DimensionIndex): ... + + +class CellDim(gtx.DimensionIndex): ... + + +class VertexDim(gtx.DimensionIndex): ... + + +class LsqUnkDim(gtx.LocalDimensionIndex): ... + + +class E2CDim(gtx.LocalDimensionIndex): ... + + +class E2VDim(gtx.LocalDimensionIndex): ... + + +class C2EDim(gtx.LocalDimensionIndex): ... + + +class V2CDim(gtx.LocalDimensionIndex): ... + + +class C2VDim(gtx.LocalDimensionIndex): ... + + +class V2EDim(gtx.LocalDimensionIndex): ... + + +class V2E2VDim(gtx.LocalDimensionIndex): ... + + +class E2C2VDim(gtx.LocalDimensionIndex): ... + + +class C2E2CODim(gtx.LocalDimensionIndex): ... + + +class E2C2EODim(gtx.LocalDimensionIndex): ... + + +class E2C2EDim(gtx.LocalDimensionIndex): ... + + +class C2E2CDim(gtx.LocalDimensionIndex): ... + + +class C2E2C2EDim(gtx.LocalDimensionIndex): ... + + +class C2E2C2E2CDim(gtx.LocalDimensionIndex): ... + + +class E2C(gtx.NeighborConnectivity[EdgeDim, CellDim]): + Local: typing.TypeAlias = E2CDim # noqa: UP040 + + +class C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + Local: typing.TypeAlias = C2EDim # noqa: UP040 + + +class V2C(gtx.NeighborConnectivity[VertexDim, CellDim]): + Local: typing.TypeAlias = V2CDim # noqa: UP040 + + +class C2V(gtx.NeighborConnectivity[CellDim, VertexDim]): + Local: typing.TypeAlias = C2VDim # noqa: UP040 + + +class V2E(gtx.NeighborConnectivity[VertexDim, EdgeDim]): + Local: typing.TypeAlias = V2EDim # noqa: UP040 + + +class E2V(gtx.NeighborConnectivity[EdgeDim, VertexDim]): + Local: typing.TypeAlias = E2VDim # noqa: UP040 + + +class E2C2V(gtx.NeighborConnectivity[EdgeDim, VertexDim]): + Local: typing.TypeAlias = E2C2VDim # noqa: UP040 + + +class C2E2CO(gtx.NeighborConnectivity[CellDim, CellDim]): + Local: typing.TypeAlias = C2E2CODim # noqa: UP040 + + +class E2C2EO(gtx.NeighborConnectivity[EdgeDim, EdgeDim]): + Local: typing.TypeAlias = E2C2EODim # noqa: UP040 + + +class E2C2E(gtx.NeighborConnectivity[EdgeDim, EdgeDim]): + Local: typing.TypeAlias = E2C2EDim # noqa: UP040 + + +class C2E2C(gtx.NeighborConnectivity[CellDim, CellDim]): + Local: typing.TypeAlias = C2E2CDim # noqa: UP040 + + +class C2E2C2E(gtx.NeighborConnectivity[CellDim, EdgeDim]): + Local: typing.TypeAlias = C2E2C2EDim # noqa: UP040 + + +class C2E2C2E2C(gtx.NeighborConnectivity[CellDim, CellDim]): + Local: typing.TypeAlias = C2E2C2E2CDim # noqa: UP040 + + +class V2E2V(gtx.NeighborConnectivity[VertexDim, VertexDim]): + Local: typing.TypeAlias = V2E2VDim # noqa: UP040 def horizontal_dims() -> Iterator[gtx.Dimension]: From 1c264ee372fea1d13d10656597b4cf28b0927a83 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 10:15:56 +0200 Subject: [PATCH 3/6] fix: test for dimensions with isinstance(..., DimensionMeta) `gtx.Dimension` is now a PEP 695 alias for `type[DimensionIndex]`, so `isinstance(d, gtx.Dimension)` raises TypeError. That made `dimension.horizontal_dims()` fail on first call, and with it everything built on it, including the GHEX domain descriptors in `mpi_decomposition`. `DimensionMeta` is not exported from `gt4py.next`, so it is reached through `gt4py.next.common`, which these modules already import. The helpers yield 3 horizontal, 2 vertical (`KDim`, `Staggered[KDim]`) and 15 local dimensions, as before. --- model/common/src/icon4py/model/common/dimension.py | 7 ++++--- model/common/src/icon4py/model/common/states/factory.py | 4 ++-- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/model/common/src/icon4py/model/common/dimension.py b/model/common/src/icon4py/model/common/dimension.py index 3fdead524e..0489b4c66f 100644 --- a/model/common/src/icon4py/model/common/dimension.py +++ b/model/common/src/icon4py/model/common/dimension.py @@ -9,6 +9,7 @@ from collections.abc import Iterator import gt4py.next as gtx +from gt4py.next import common as gtx_common class KDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... @@ -132,7 +133,7 @@ def horizontal_dims() -> Iterator[gtx.Dimension]: tuple( d for d in globals().values() - if isinstance(d, gtx.Dimension) and d.kind == gtx.DimensionKind.HORIZONTAL + if isinstance(d, gtx_common.DimensionMeta) and d.kind == gtx.DimensionKind.HORIZONTAL ) ) @@ -144,7 +145,7 @@ def non_horizontal_dims() -> Iterator[gtx.Dimension]: def local_dims() -> Iterator[gtx.Dimension]: for d in globals().values(): - if isinstance(d, gtx.Dimension) and d.kind == gtx.DimensionKind.LOCAL: + if isinstance(d, gtx_common.DimensionMeta) and d.kind == gtx.DimensionKind.LOCAL: yield d @@ -153,6 +154,6 @@ def vertical_dims() -> Iterator[gtx.Dimension]: tuple( d for d in globals().values() - if isinstance(d, gtx.Dimension) and d.kind == gtx.DimensionKind.VERTICAL + if isinstance(d, gtx_common.DimensionMeta) and d.kind == gtx.DimensionKind.VERTICAL ) ) diff --git a/model/common/src/icon4py/model/common/states/factory.py b/model/common/src/icon4py/model/common/states/factory.py index 0837d5a29a..4cc7d22053 100644 --- a/model/common/src/icon4py/model/common/states/factory.py +++ b/model/common/src/icon4py/model/common/states/factory.py @@ -454,7 +454,7 @@ def _get_offset_providers(self, grid: icon_grid.IconGrid) -> dict[str, gtx.Field vertical_offsets = { k: v for k, v in grid.connectivities.items() - if isinstance(v, gtx.Dimension) and v.kind == gtx.DimensionKind.VERTICAL + if isinstance(v, gtx_common.DimensionMeta) and v.kind == gtx.DimensionKind.VERTICAL } offset_providers.update(vertical_offsets) # used for different compute backend in function call @@ -546,7 +546,7 @@ def _get_offset_providers(self, grid: icon_grid.IconGrid) -> dict[str, gtx.Field vertical_offsets = { k: v for k, v in grid.connectivities.items() - if isinstance(v, gtx.Dimension) and v.kind == gtx.DimensionKind.VERTICAL + if isinstance(v, gtx_common.DimensionMeta) and v.kind == gtx.DimensionKind.VERTICAL } offset_providers.update(vertical_offsets) return offset_providers From b0e4a57d1ea05590150904f973ae48ae10e15078 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 10:16:05 +0200 Subject: [PATCH 4/6] refactor: key offset providers by connectivity class #2910 keys offset providers by the connectivity declaration and rejects string keys at every program entry point, and removes `FieldOffset`. `Grid` passes its `connectivities` straight through as the provider, so it is now keyed by the connectivity class: - `Grid.connectivities` is `Mapping[type[NeighborConnectivity], NeighborTable]` and `get_connectivity` takes the class. `construct_connectivity` reads `domain`, `codomain` and `local_dimension_of` from the declaration instead of `FieldOffset.source/target`. - The halo constructor and the grid manager's neighbor tables are class-keyed; the halo's `str | FieldOffset` normalisation is gone, since both callers now pass classes. - Field providers take `connectivities={"e2c": dims.E2C}`, the connectivity itself, where they took its local dimension and read its name. - The stencil-test connectivity view is class-keyed and no longer needs a cast over a name-keyed mapping; its tests check the class-keyed contract. - String provider keys and `get_connectivity("X")` calls become `dims.X`, and the `Mapping[gtx.FieldOffset, np.ndarray]` annotations of the reference functions become `Mapping[type[gtx.NeighborConnectivity], np.ndarray]`. --- bindings/src/icon4py/bindings/common.py | 6 +-- .../stencil_tests/test_apply_nabla2_to_w.py | 2 +- ...ate_horizontal_gradients_for_turbulence.py | 2 +- .../test_calculate_nabla2_for_w.py | 4 +- .../test_calculate_nabla2_for_z.py | 2 +- .../test_calculate_nabla2_of_theta.py | 2 +- .../stencil_tests/test_calculate_nabla4.py | 2 +- ...fusion_nabla_of_theta_over_steep_points.py | 2 +- .../integration_tests/test_solve_nonhydro.py | 24 ++++----- .../test_velocity_advection.py | 30 ++++++------ ...lysis_increments_from_data_assimilation.py | 2 +- ...or_normal_wind_tendency_approaching_cfl.py | 2 +- ...l_wind_derivative_to_divergence_damping.py | 2 +- .../test_apply_rayleigh_damping_mechanism.py | 2 +- ...st_compute_avg_vn_and_graddiv_vn_and_vt.py | 2 +- ...t_compute_contravariant_correction_of_w.py | 2 +- ...iant_correction_of_w_for_lower_boundary.py | 2 +- ...e_divergence_of_fluxes_of_rho_and_theta.py | 2 +- ...est_compute_dwdz_for_divergence_damping.py | 2 +- ...compute_explicit_part_for_rho_and_exner.py | 2 +- ...rom_advection_and_vertical_wind_density.py | 2 +- ...d_speed_and_vertical_wind_times_density.py | 2 +- .../test_compute_graddiv2_of_vn.py | 2 +- ...e_horizontal_advection_of_rho_and_theta.py | 6 +-- ..._of_exner_pressure_for_flat_coordinates.py | 2 +- ...t_of_exner_pressure_for_multiple_levels.py | 2 +- ..._exner_pressure_for_nonflat_coordinates.py | 2 +- ...rizontal_velocity_quantities_and_fluxes.py | 2 +- ...est_compute_hydrostatic_correction_term.py | 2 +- ...ute_results_for_thermodynamic_variables.py | 2 +- ...test_compute_solver_coefficients_matrix.py | 2 +- ...ues_and_pressure_gradient_and_update_vn.py | 2 +- ...tial_temperatures_and_pressure_gradient.py | 2 +- ...t_extrapolate_temporally_exner_pressure.py | 2 +- ...tion_for_w_and_contravariant_correction.py | 2 +- ...diagonal_matrix_for_w_back_substitution.py | 2 +- ...test_spatially_average_flux_or_velocity.py | 2 +- ...t_update_dynamical_exner_time_increment.py | 2 +- .../test_update_mass_volume_flux.py | 2 +- .../test_velocity_advection_terms.py | 20 ++++---- .../tmx/stencil_tests/test_diagnostics.py | 6 +-- .../test_tracer_advection.py | 2 +- .../model/common/decomposition/halo.py | 18 +++---- .../src/icon4py/model/common/grid/base.py | 15 +++--- .../src/icon4py/model/common/grid/geometry.py | 4 +- .../icon4py/model/common/grid/grid_manager.py | 12 ++--- .../src/icon4py/model/common/grid/icon.py | 15 ++++-- .../src/icon4py/model/common/grid/simple.py | 4 +- .../linear_horizontal_tracer_advection.py | 5 +- .../interpolation/interpolation_factory.py | 42 ++++++++-------- .../common/metrics/compute_zdiff_gradp.py | 2 +- .../model/common/metrics/metrics_factory.py | 14 +++--- .../icon4py/model/common/states/factory.py | 22 ++++++--- .../mpi_tests/test_mpi_decomposition.py | 2 +- .../decomposition/unit_tests/test_halo.py | 6 +-- .../grid/unit_tests/test_grid_manager.py | 49 ++++++++++--------- .../tests/common/grid/unit_tests/test_icon.py | 4 +- .../common/grid/unit_tests/test_topography.py | 4 +- .../common/grid/unit_tests/test_vertical.py | 2 +- model/common/tests/common/grid/utils.py | 4 +- ...st_edge_2_cell_vector_rbf_interpolation.py | 2 +- .../unit_tests/test_interpolation_fields.py | 4 +- ...izontal_gradients_by_green_gauss_method.py | 2 +- .../test_compute_diffusion_metrics.py | 12 ++--- .../unit_tests/test_compute_zdiff_gradp.py | 4 +- .../metrics/unit_tests/test_metric_fields.py | 16 +++--- .../unit_tests/test_reference_atmosphere.py | 2 +- .../src/icon4py/model/driver/driver_io.py | 2 +- .../icon4py/model/testing/reference_funcs.py | 16 +++--- .../icon4py/model/testing/stencil_tests.py | 23 ++++----- .../unit_tests/test_stenciltest_framework.py | 27 +++++----- 71 files changed, 254 insertions(+), 246 deletions(-) diff --git a/bindings/src/icon4py/bindings/common.py b/bindings/src/icon4py/bindings/common.py index 58725d43e8..5d52d2d3f5 100644 --- a/bindings/src/icon4py/bindings/common.py +++ b/bindings/src/icon4py/bindings/common.py @@ -164,10 +164,10 @@ def impl(_name: str, domain: gtx.Domain, dtype: gt4py_definitions.DType) -> gtx. def shrink_to_dimension( - sizes: dict[gtx.Dimension, int], tables: dict[gtx.FieldOffset, NDArray] -) -> dict[gtx.FieldOffset, NDArray]: + sizes: dict[gtx.Dimension, int], tables: dict[type[gtx.NeighborConnectivity], NDArray] +) -> dict[type[gtx.NeighborConnectivity], NDArray]: """Shrink the neighbor tables from nproma size to the actual size of the grid.""" - return {k: v[: sizes[k.target[0]]] for k, v in tables.items()} + return {k: v[: sizes[k.domain]] for k, v in tables.items()} def add_origin(xp: ModuleType, table: NDArray) -> NDArray: diff --git a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_apply_nabla2_to_w.py b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_apply_nabla2_to_w.py index 341481dc83..5856e26812 100644 --- a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_apply_nabla2_to_w.py +++ b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_apply_nabla2_to_w.py @@ -20,7 +20,7 @@ def apply_nabla2_to_w_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], area: np.ndarray, z_nabla2_c: np.ndarray, geofac_n2s: np.ndarray, diff --git a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_horizontal_gradients_for_turbulence.py b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_horizontal_gradients_for_turbulence.py index 0739e9997f..f20497e5e4 100644 --- a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_horizontal_gradients_for_turbulence.py +++ b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_horizontal_gradients_for_turbulence.py @@ -21,7 +21,7 @@ def calculate_horizontal_gradients_for_turbulence_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], w: np.ndarray, geofac_grg_x: np.ndarray, geofac_grg_y: np.ndarray, diff --git a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_w.py b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_w.py index 4e2496bd71..54399a84b0 100644 --- a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_w.py +++ b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_w.py @@ -20,7 +20,9 @@ def calculate_nabla2_for_w_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], w: np.ndarray, geofac_n2s: np.ndarray + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], + w: np.ndarray, + geofac_n2s: np.ndarray, ) -> np.ndarray: c2e2cO = connectivities[dims.C2E2CO] geofac_n2s = np.expand_dims(geofac_n2s, axis=-1) diff --git a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_z.py b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_z.py index 78fc5e0528..96ab28c9ea 100644 --- a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_z.py +++ b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_for_z.py @@ -21,7 +21,7 @@ def calculate_nabla2_for_z_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], kh_smag_e: np.ndarray, inv_dual_edge_length: np.ndarray, theta_v: np.ndarray, diff --git a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_of_theta.py b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_of_theta.py index c37252f38d..def3712bfd 100644 --- a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_of_theta.py +++ b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla2_of_theta.py @@ -20,7 +20,7 @@ def calculate_nabla2_of_theta_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], z_nabla2_e: np.ndarray, geofac_div: np.ndarray, ) -> np.ndarray: diff --git a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla4.py b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla4.py index 5ba71bf8aa..59392e3cd5 100644 --- a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla4.py +++ b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_calculate_nabla4.py @@ -19,7 +19,7 @@ def calculate_nabla4_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], u_vert: np.ndarray, v_vert: np.ndarray, primal_normal_vert_v1: np.ndarray, diff --git a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py index 492d731a34..a94868921f 100644 --- a/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py +++ b/model/atmosphere/diffusion/tests/diffusion/stencil_tests/test_truly_horizontal_diffusion_nabla_of_theta_over_steep_points.py @@ -22,7 +22,7 @@ def truly_horizontal_diffusion_nabla_of_theta_over_steep_points_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], zd_vertoffset: np.ndarray, zd_diffcoef: np.ndarray, geofac_n2s_c: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/integration_tests/test_solve_nonhydro.py b/model/atmosphere/dycore/tests/dycore/integration_tests/test_solve_nonhydro.py index 0c4c096a19..54e13ab58a 100644 --- a/model/atmosphere/dycore/tests/dycore/integration_tests/test_solve_nonhydro.py +++ b/model/atmosphere/dycore/tests/dycore/integration_tests/test_solve_nonhydro.py @@ -1353,7 +1353,7 @@ def test_compute_rho_theta_pgrad_and_update_vn( # noqa: PLR0917 [too-many-posit vertical_start=icon_grid.num_levels - 1, vertical_end=icon_grid.num_levels, offset_provider={ - "E2C": icon_grid.get_connectivity("E2C"), + dims.E2C: icon_grid.get_connectivity(dims.E2C), }, ) lowest_level = icon_grid.num_levels - 1 @@ -1414,9 +1414,9 @@ def test_compute_rho_theta_pgrad_and_update_vn( # noqa: PLR0917 [too-many-posit vertical_start=gtx.int32(0), vertical_end=gtx.int32(icon_grid.num_levels), offset_provider={ - "C2E2CO": icon_grid.get_connectivity("C2E2CO"), - "E2C": icon_grid.get_connectivity("E2C"), - "E2C2EO": icon_grid.get_connectivity("E2C2EO"), + dims.C2E2CO: icon_grid.get_connectivity(dims.C2E2CO), + dims.E2C: icon_grid.get_connectivity(dims.E2C), + dims.E2C2EO: icon_grid.get_connectivity(dims.E2C2EO), }, ) @@ -1575,9 +1575,9 @@ def test_apply_divergence_damping_and_update_vn( # noqa: PLR0917 [too-many-posi vertical_start=gtx.int32(0), vertical_end=gtx.int32(icon_grid.num_levels), offset_provider={ - "C2E2CO": icon_grid.get_connectivity("C2E2CO"), - "E2C": icon_grid.get_connectivity("E2C"), - "E2C2EO": icon_grid.get_connectivity("E2C2EO"), + dims.C2E2CO: icon_grid.get_connectivity(dims.C2E2CO), + dims.E2C: icon_grid.get_connectivity(dims.E2C), + dims.E2C2EO: icon_grid.get_connectivity(dims.E2C2EO), }, ) @@ -1687,8 +1687,8 @@ def test_compute_horizontal_velocity_quantities_and_fluxes( # noqa: PLR0917 [to vertical_start=0, vertical_end=icon_grid.num_levels + 1, offset_provider={ - "E2C2EO": icon_grid.get_connectivity("E2C2EO"), - "E2C2E": icon_grid.get_connectivity("E2C2E"), + dims.E2C2EO: icon_grid.get_connectivity(dims.E2C2EO), + dims.E2C2E: icon_grid.get_connectivity(dims.E2C2E), }, ) @@ -1825,7 +1825,7 @@ def test_compute_averaged_vn_and_fluxes( # noqa: PLR0917 [too-many-positional-a vertical_start=0, vertical_end=icon_grid.num_levels, offset_provider={ - "E2C2EO": icon_grid.get_connectivity("E2C2EO"), + dims.E2C2EO: icon_grid.get_connectivity(dims.E2C2EO), }, ) @@ -1949,7 +1949,7 @@ def test_vertically_implicit_solver_at_predictor_step( # noqa: PLR0917 [too-man end_cell_halo = icon_grid.end_index(cell_domain(h_grid.Zone.HALO)) offset_provider = { - "C2E": icon_grid.get_connectivity("C2E"), + dims.C2E: icon_grid.get_connectivity(dims.C2E), } vertically_implicit_dycore_solver.vertically_implicit_solver_at_predictor_step.with_backend( @@ -2139,7 +2139,7 @@ def test_vertically_implicit_solver_at_corrector_step( # noqa: PLR0917 [too-man end_cell_local = icon_grid.end_index(cell_domain(h_grid.Zone.LOCAL)) offset_provider = { - "C2E": icon_grid.get_connectivity("C2E"), + dims.C2E: icon_grid.get_connectivity(dims.C2E), } vertically_implicit_dycore_solver.vertically_implicit_solver_at_corrector_step.with_backend( diff --git a/model/atmosphere/dycore/tests/dycore/integration_tests/test_velocity_advection.py b/model/atmosphere/dycore/tests/dycore/integration_tests/test_velocity_advection.py index 14c05cdb08..63eac7d5a7 100644 --- a/model/atmosphere/dycore/tests/dycore/integration_tests/test_velocity_advection.py +++ b/model/atmosphere/dycore/tests/dycore/integration_tests/test_velocity_advection.py @@ -231,14 +231,14 @@ def test_velocity_predictor_step( # noqa: PLR0917 [too-many-positional-argument vertical_start=gtx.int32(0), vertical_end=icon_grid.num_levels, offset_provider={ - "C2E": icon_grid.get_connectivity("C2E"), - "C2E2CO": icon_grid.get_connectivity("C2E2CO"), - "E2C": icon_grid.get_connectivity("E2C"), - "E2C2E": icon_grid.get_connectivity("E2C2E"), - "E2C2EO": icon_grid.get_connectivity("E2C2EO"), - "E2V": icon_grid.get_connectivity("E2V"), - "V2C": icon_grid.get_connectivity("V2C"), - "V2E": icon_grid.get_connectivity("V2E"), + dims.C2E: icon_grid.get_connectivity(dims.C2E), + dims.C2E2CO: icon_grid.get_connectivity(dims.C2E2CO), + dims.E2C: icon_grid.get_connectivity(dims.E2C), + dims.E2C2E: icon_grid.get_connectivity(dims.E2C2E), + dims.E2C2EO: icon_grid.get_connectivity(dims.E2C2EO), + dims.E2V: icon_grid.get_connectivity(dims.E2V), + dims.V2C: icon_grid.get_connectivity(dims.V2C), + dims.V2E: icon_grid.get_connectivity(dims.V2E), }, ) solve_nonhydro._update_max_vertical_cfl( @@ -458,13 +458,13 @@ def test_velocity_corrector_step( # noqa: PLR0917 [too-many-positional-argument vertical_start=gtx.int32(0), vertical_end=icon_grid.num_levels, offset_provider={ - "C2E": icon_grid.get_connectivity("C2E"), - "C2E2CO": icon_grid.get_connectivity("C2E2CO"), - "E2C": icon_grid.get_connectivity("E2C"), - "E2C2EO": icon_grid.get_connectivity("E2C2EO"), - "E2V": icon_grid.get_connectivity("E2V"), - "V2C": icon_grid.get_connectivity("V2C"), - "V2E": icon_grid.get_connectivity("V2E"), + dims.C2E: icon_grid.get_connectivity(dims.C2E), + dims.C2E2CO: icon_grid.get_connectivity(dims.C2E2CO), + dims.E2C: icon_grid.get_connectivity(dims.E2C), + dims.E2C2EO: icon_grid.get_connectivity(dims.E2C2EO), + dims.E2V: icon_grid.get_connectivity(dims.E2V), + dims.V2C: icon_grid.get_connectivity(dims.V2C), + dims.V2E: icon_grid.get_connectivity(dims.V2E), }, ) solve_nonhydro._update_max_vertical_cfl( diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_analysis_increments_from_data_assimilation.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_analysis_increments_from_data_assimilation.py index f10107b4a4..a75847048d 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_analysis_increments_from_data_assimilation.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_analysis_increments_from_data_assimilation.py @@ -23,7 +23,7 @@ def add_analysis_increments_from_data_assimilation_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], z_rho_expl: np.ndarray, rho_incr: np.ndarray, z_exner_expl: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_extra_diffusion_for_normal_wind_tendency_approaching_cfl.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_extra_diffusion_for_normal_wind_tendency_approaching_cfl.py index 47a4ff99a6..361862924d 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_extra_diffusion_for_normal_wind_tendency_approaching_cfl.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_extra_diffusion_for_normal_wind_tendency_approaching_cfl.py @@ -24,7 +24,7 @@ def add_extra_diffusion_for_normal_wind_tendency_approaching_cfl_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], levelmask: np.ndarray, c_lin_e: np.ndarray, z_w_con_c_full: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_vertical_wind_derivative_to_divergence_damping.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_vertical_wind_derivative_to_divergence_damping.py index 3ee43babdd..f43abfd597 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_vertical_wind_derivative_to_divergence_damping.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_add_vertical_wind_derivative_to_divergence_damping.py @@ -23,7 +23,7 @@ def add_vertical_wind_derivative_to_divergence_damping_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], hmask_dd3d: np.ndarray, scalfac_dd3d: np.ndarray, inv_dual_edge_length: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_apply_rayleigh_damping_mechanism.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_apply_rayleigh_damping_mechanism.py index 52e55b1f74..356afa89fe 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_apply_rayleigh_damping_mechanism.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_apply_rayleigh_damping_mechanism.py @@ -23,7 +23,7 @@ def apply_rayleigh_damping_mechanism_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], z_raylfac: np.ndarray, w: np.ndarray, ) -> np.ndarray: diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_avg_vn_and_graddiv_vn_and_vt.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_avg_vn_and_graddiv_vn_and_vt.py index c5b6711bef..7ca2440fe7 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_avg_vn_and_graddiv_vn_and_vt.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_avg_vn_and_graddiv_vn_and_vt.py @@ -23,7 +23,7 @@ def compute_avg_vn_and_graddiv_vn_and_vt_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], e_flx_avg: np.ndarray, vn: np.ndarray, geofac_grdiv: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w.py index 81cce49ec8..801119a88c 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w.py @@ -23,7 +23,7 @@ def compute_contravariant_correction_of_w_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], e_bln_c_s: np.ndarray, z_w_concorr_me: np.ndarray, wgtfac_c: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w_for_lower_boundary.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w_for_lower_boundary.py index 6de11f0732..d19e703d05 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w_for_lower_boundary.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_contravariant_correction_of_w_for_lower_boundary.py @@ -23,7 +23,7 @@ def compute_contravariant_correction_of_w_for_lower_boundary_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], e_bln_c_s: np.ndarray, z_w_concorr_me: np.ndarray, wgtfacq_c: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_divergence_of_fluxes_of_rho_and_theta.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_divergence_of_fluxes_of_rho_and_theta.py index 15ad8c0757..4760731da3 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_divergence_of_fluxes_of_rho_and_theta.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_divergence_of_fluxes_of_rho_and_theta.py @@ -22,7 +22,7 @@ def compute_divergence_of_fluxes_of_rho_and_theta_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], geofac_div: np.ndarray, mass_flux_at_edges_on_model_levels: np.ndarray, theta_v_flux_at_edges_on_model_levels: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_dwdz_for_divergence_damping.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_dwdz_for_divergence_damping.py index bf5abbc205..a4a681b4f7 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_dwdz_for_divergence_damping.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_dwdz_for_divergence_damping.py @@ -22,7 +22,7 @@ def compute_dwdz_for_divergence_damping_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], inv_ddqz_z_full: np.ndarray, w: np.ndarray, w_concorr_c: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_part_for_rho_and_exner.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_part_for_rho_and_exner.py index fbb06d7a0b..201c385d14 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_part_for_rho_and_exner.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_part_for_rho_and_exner.py @@ -23,7 +23,7 @@ def compute_explicit_part_for_rho_and_exner_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], rho_nnow: np.ndarray, inv_ddqz_z_full: np.ndarray, z_flxdiv_mass: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_from_advection_and_vertical_wind_density.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_from_advection_and_vertical_wind_density.py index 6828918839..f397921895 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_from_advection_and_vertical_wind_density.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_from_advection_and_vertical_wind_density.py @@ -23,7 +23,7 @@ def compute_explicit_vertical_wind_from_advection_and_vertical_wind_density_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], w_nnow: np.ndarray, ddt_w_adv_ntl1: np.ndarray, ddt_w_adv_ntl2: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_speed_and_vertical_wind_times_density.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_speed_and_vertical_wind_times_density.py index 710d2000f5..3ef667977a 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_speed_and_vertical_wind_times_density.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_explicit_vertical_wind_speed_and_vertical_wind_times_density.py @@ -24,7 +24,7 @@ def compute_explicit_vertical_wind_speed_and_vertical_wind_times_density_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], w_nnow: np.ndarray, ddt_w_adv_ntl1: np.ndarray, z_th_ddz_exner_c: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_graddiv2_of_vn.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_graddiv2_of_vn.py index 5ad1d46bec..ce19569a37 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_graddiv2_of_vn.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_graddiv2_of_vn.py @@ -21,7 +21,7 @@ def compute_graddiv2_of_vn_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], geofac_grdiv: np.ndarray, z_graddiv_vn: np.ndarray, ) -> np.ndarray: diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_advection_of_rho_and_theta.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_advection_of_rho_and_theta.py index af5e59bae9..2846ba0796 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_advection_of_rho_and_theta.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_advection_of_rho_and_theta.py @@ -23,7 +23,7 @@ # TODO(): copied from `test_mo_math_gradients_grad_green_gauss_cell_dsl_numpy`. delete that test? def mo_math_gradients_grad_green_gauss_cell_dsl_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], p_ccpr1: np.ndarray, p_ccpr2: np.ndarray, geofac_grg_x: np.ndarray, @@ -97,7 +97,7 @@ def compute_btraj_numpy( def sten_16_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], p_vn: np.ndarray, rho_ref_me: np.ndarray, theta_ref_me: np.ndarray, @@ -148,7 +148,7 @@ def sten_16_numpy( def compute_horizontal_advection_of_rho_and_theta_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], p_vn: np.ndarray, p_vt: np.ndarray, pos_on_tplane_e_1: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_flat_coordinates.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_flat_coordinates.py index dbaa79cf72..8a94a16631 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_flat_coordinates.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_flat_coordinates.py @@ -23,7 +23,7 @@ def compute_horizontal_gradient_of_exner_pressure_for_flat_coordinates_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], inv_dual_edge_length: np.ndarray, z_exner_ex_pr: np.ndarray, ) -> np.ndarray: diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py index 04209c4610..cbb95709a1 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_multiple_levels.py @@ -23,7 +23,7 @@ def compute_horizontal_gradient_of_exner_pressure_for_multiple_levels_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], inv_dual_edge_length: np.ndarray, z_exner_ex_pr: np.ndarray, zdiff_gradp: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_nonflat_coordinates.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_nonflat_coordinates.py index 39bcd70c42..a86964be06 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_nonflat_coordinates.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_gradient_of_exner_pressure_for_nonflat_coordinates.py @@ -24,7 +24,7 @@ def compute_horizontal_gradient_of_exner_pressure_for_nonflat_coordinates_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], inv_dual_edge_length: np.ndarray, z_exner_ex_pr: np.ndarray, ddxn_z_full: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_velocity_quantities_and_fluxes.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_velocity_quantities_and_fluxes.py index 3cdffcc526..a7f4090d1a 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_velocity_quantities_and_fluxes.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_horizontal_velocity_quantities_and_fluxes.py @@ -33,7 +33,7 @@ def compute_vt_vn_on_half_levels_and_kinetic_energy_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], vn: np.ndarray, tangential_wind: np.ndarray, vn_on_half_levels: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_hydrostatic_correction_term.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_hydrostatic_correction_term.py index ad3bacaa13..9308bfe051 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_hydrostatic_correction_term.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_hydrostatic_correction_term.py @@ -23,7 +23,7 @@ def compute_hydrostatic_correction_term_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], theta_v: np.ndarray, ikoffset: np.ndarray, zdiff_gradp: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_results_for_thermodynamic_variables.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_results_for_thermodynamic_variables.py index 1514880444..a7d3b80dba 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_results_for_thermodynamic_variables.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_results_for_thermodynamic_variables.py @@ -23,7 +23,7 @@ def compute_results_for_thermodynamic_variables_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], z_rho_expl: np.ndarray, vwind_impl_wgt: np.ndarray, inv_ddqz_z_full: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_solver_coefficients_matrix.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_solver_coefficients_matrix.py index 2adfaeb90a..e23464b4dc 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_solver_coefficients_matrix.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_solver_coefficients_matrix.py @@ -23,7 +23,7 @@ def compute_solver_coefficients_matrix_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], exner_nnow: np.ndarray, rho_nnow: np.ndarray, theta_v_nnow: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_theta_rho_face_values_and_pressure_gradient_and_update_vn.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_theta_rho_face_values_and_pressure_gradient_and_update_vn.py index ec572ad5f9..03351425fd 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_theta_rho_face_values_and_pressure_gradient_and_update_vn.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_theta_rho_face_values_and_pressure_gradient_and_update_vn.py @@ -24,7 +24,7 @@ def compute_theta_rho_face_value_by_miura_scheme_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], vn: np.ndarray, tangential_wind: np.ndarray, pos_on_tplane_e_x: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_virtual_potential_temperatures_and_pressure_gradient.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_virtual_potential_temperatures_and_pressure_gradient.py index 2189d773e0..2ebfe42cd1 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_virtual_potential_temperatures_and_pressure_gradient.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_compute_virtual_potential_temperatures_and_pressure_gradient.py @@ -24,7 +24,7 @@ def compute_virtual_potential_temperatures_and_pressure_gradient_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], wgtfac_c: np.ndarray, z_rth_pr_2: np.ndarray, theta_v: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_extrapolate_temporally_exner_pressure.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_extrapolate_temporally_exner_pressure.py index 86756f43db..e18956c36e 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_extrapolate_temporally_exner_pressure.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_extrapolate_temporally_exner_pressure.py @@ -23,7 +23,7 @@ def extrapolate_temporally_exner_pressure_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], exner: np.ndarray, exner_ref_mc: np.ndarray, exner_pr: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_set_lower_boundary_condition_for_w_and_contravariant_correction.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_set_lower_boundary_condition_for_w_and_contravariant_correction.py index 59e3e67ec9..910ec8af2d 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_set_lower_boundary_condition_for_w_and_contravariant_correction.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_set_lower_boundary_condition_for_w_and_contravariant_correction.py @@ -23,7 +23,7 @@ def set_lower_boundary_condition_for_w_and_contravariant_correction_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], w_concorr_c: np.ndarray, z_contr_w_fl_l: np.ndarray, ) -> tuple[np.ndarray, np.ndarray]: diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_solve_tridiagonal_matrix_for_w_back_substitution.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_solve_tridiagonal_matrix_for_w_back_substitution.py index 6b677ad2e9..87f1e164ad 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_solve_tridiagonal_matrix_for_w_back_substitution.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_solve_tridiagonal_matrix_for_w_back_substitution.py @@ -23,7 +23,7 @@ def solve_tridiagonal_matrix_for_w_back_substitution_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], z_q: np.ndarray, w: np.ndarray, ) -> np.ndarray: diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_spatially_average_flux_or_velocity.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_spatially_average_flux_or_velocity.py index 0ab3618eb8..88cf0a15cd 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_spatially_average_flux_or_velocity.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_spatially_average_flux_or_velocity.py @@ -24,7 +24,7 @@ def spatially_average_flux_or_velocity_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], e_flx_avg: np.ndarray, flux_or_velocity: np.ndarray, ) -> np.ndarray: diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_dynamical_exner_time_increment.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_dynamical_exner_time_increment.py index b94ed197f2..5b619fb7f5 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_dynamical_exner_time_increment.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_dynamical_exner_time_increment.py @@ -24,7 +24,7 @@ def update_dynamical_exner_time_increment_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], exner: np.ndarray, ddt_exner_phy: np.ndarray, exner_dyn_incr: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_mass_volume_flux.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_mass_volume_flux.py index b139fe5319..0d1bcb1957 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_mass_volume_flux.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_update_mass_volume_flux.py @@ -21,7 +21,7 @@ def update_mass_volume_flux_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], z_contr_w_fl_l: np.ndarray, rho_ic: np.ndarray, vwind_impl_wgt: np.ndarray, diff --git a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_velocity_advection_terms.py b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_velocity_advection_terms.py index 2ddbbc2a78..52cf7e745c 100644 --- a/model/atmosphere/dycore/tests/dycore/stencil_tests/test_velocity_advection_terms.py +++ b/model/atmosphere/dycore/tests/dycore/stencil_tests/test_velocity_advection_terms.py @@ -72,7 +72,7 @@ def extrapolate_to_surface_numpy(vn: np.ndarray, wgtfacq_e: np.ndarray) -> np.nd def compute_diagnostics_from_normal_wind_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], tangential_wind_on_half_levels: np.ndarray, vn: np.ndarray, rbf_vec_coeff_e: np.ndarray, @@ -113,7 +113,7 @@ def compute_diagnostics_from_normal_wind_numpy( def interpolate_contravariant_correction_to_cells_on_half_levels_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], contravariant_correction_at_edges_on_model_levels: np.ndarray, e_bln_c_s: np.ndarray, wgtfac_c: np.ndarray, @@ -197,7 +197,7 @@ def compute_maximum_cfl_and_clip_contravariant_vertical_velocity_numpy( def compute_horizontal_advection_of_w_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], w: np.ndarray, tangential_wind_on_half_levels: np.ndarray, vn_on_half_levels: np.ndarray, @@ -223,7 +223,7 @@ def compute_horizontal_advection_of_w_numpy( def add_extra_diffusion_for_w_approaching_cfl_wihtout_levmask_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], cfl_clipping: np.ndarray, owner_mask: np.ndarray, contravariant_corrected_w_at_cells_on_half_levels: np.ndarray, @@ -287,7 +287,7 @@ def compute_advective_vertical_wind_tendency_numpy( def compute_advective_vertical_wind_tendency_and_apply_diffusion_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], vertical_wind_advective_tendency: np.ndarray, w: np.ndarray, horizontal_advection_of_w_at_edges_on_half_levels: np.ndarray, @@ -348,7 +348,7 @@ def compute_advective_vertical_wind_tendency_and_apply_diffusion_numpy( def _compute_advective_normal_wind_tendency_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], horizontal_kinetic_energy_at_edges_on_model_levels: np.ndarray, coeff_gradekin: np.ndarray, horizontal_kinetic_energy_at_cells_on_model_levels: np.ndarray, @@ -388,7 +388,7 @@ def _compute_advective_normal_wind_tendency_numpy( def _add_extra_diffusion_for_normal_wind_tendency_approaching_cfl_without_levelmask_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], c_lin_e: np.ndarray, contravariant_corrected_w_at_cells_on_model_levels: np.ndarray, ddqz_z_full_e: np.ndarray, @@ -460,7 +460,7 @@ def _add_extra_diffusion_for_normal_wind_tendency_approaching_cfl_without_levelm def compute_advection_in_horizontal_momentum_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], vn: np.ndarray, horizontal_kinetic_energy_at_edges_on_model_levels: np.ndarray, tangential_wind: np.ndarray, @@ -540,7 +540,7 @@ def _restore_outside( def compute_interpolated_horizontal_advection_of_w_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], e_bln_c_s: np.ndarray, horizontal_advection_of_w_at_edges_on_half_levels: np.ndarray, **kwargs: Any, @@ -555,7 +555,7 @@ def compute_interpolated_horizontal_advection_of_w_numpy( def compute_extra_diffusion_for_w_numpy( *, - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], contravariant_corrected_w_at_cells_on_half_levels: np.ndarray, ddqz_z_half: np.ndarray, area: np.ndarray, diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_diagnostics.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_diagnostics.py index 7cc5a9b5f4..d4f099c2df 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_diagnostics.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_diagnostics.py @@ -328,7 +328,7 @@ def input_data( def interpolate_cell_field_to_edge_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], in_field: np.ndarray, coeff: np.ndarray, ) -> np.ndarray: @@ -338,7 +338,7 @@ def interpolate_cell_field_to_edge_numpy( def compute_shear_and_div_of_stress_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], *, u_vert: np.ndarray, v_vert: np.ndarray, @@ -412,7 +412,7 @@ def compute_shear_and_div_of_stress_numpy( def interpolate_edge_field_to_cell_half_levels_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], interpolant: np.ndarray, e_bln_c_s: np.ndarray, wgtfac_c: np.ndarray, diff --git a/model/atmosphere/tracer_advection/tests/tracer_advection/integration_tests/test_tracer_advection.py b/model/atmosphere/tracer_advection/tests/tracer_advection/integration_tests/test_tracer_advection.py index 1c90d58b91..0c491d788b 100644 --- a/model/atmosphere/tracer_advection/tests/tracer_advection/integration_tests/test_tracer_advection.py +++ b/model/atmosphere/tracer_advection/tests/tracer_advection/integration_tests/test_tracer_advection.py @@ -129,7 +129,7 @@ def test_tracer_advection_run_single_step( # noqa: PLR0917 [too-many-positional cell_center_y=geometry.get(geometry_attrs.CELL_CENTER_Y).asnumpy(), cell_lat=geometry.get(geometry_attrs.CELL_LAT).asnumpy(), cell_lon=geometry.get(geometry_attrs.CELL_LON).asnumpy(), - c2e2c=icon_grid.connectivities["C2E2C"].asnumpy(), + c2e2c=icon_grid.connectivities[dims.C2E2C].asnumpy(), cell_owner_mask=grid_savepoint.c_owner_mask().asnumpy(), domain_length=geometry.grid.grid_params.domain_length, domain_height=geometry.grid.grid_params.domain_height, diff --git a/model/common/src/icon4py/model/common/decomposition/halo.py b/model/common/src/icon4py/model/common/decomposition/halo.py index 82daf1834f..5c7f359013 100644 --- a/model/common/src/icon4py/model/common/decomposition/halo.py +++ b/model/common/src/icon4py/model/common/decomposition/halo.py @@ -61,7 +61,7 @@ class IconLikeHaloConstructor(HaloConstructor): def __init__( self, process_props: defs.ProcessProperties, - connectivities: dict[gtx.FieldOffset | str, data_alloc.NDArray], + connectivities: dict[type[gtx.NeighborConnectivity], data_alloc.NDArray], allocator: gtx_typing.Allocator | None = None, ): """ @@ -73,13 +73,9 @@ def __init__( """ self._xp = data_alloc.import_array_ns(allocator) self._process_props = process_props - self._connectivities = {self._value(k): v for k, v in connectivities.items()} + self._connectivities = connectivities self._assert_all_neighbor_tables() - @staticmethod - def _value(k: gtx.FieldOffset | str) -> str: - return str(k.value) if isinstance(k, gtx.FieldOffset) else k - def _validate_mapping(self, cell_to_rank_mapping: data_alloc.NDArray) -> None: # validate the distribution mapping: num_cells = self._connectivity(dims.C2E2C).shape[0] @@ -108,13 +104,13 @@ def _assert_all_neighbor_tables(self) -> None: dims.V2E, ] for d in relevant_dimension: - assert d.value in self._connectivities, ( + assert d in self._connectivities, ( f"Table for {d} is missing from the neighbor table array." ) - def _connectivity(self, offset: gtx.FieldOffset | str) -> data_alloc.NDArray: + def _connectivity(self, offset: type[gtx.NeighborConnectivity]) -> data_alloc.NDArray: try: - return self._connectivities[self._value(offset)] + return self._connectivities[offset] except KeyError as err: raise exceptions.MissingConnectivityError( f"Connectivity for offset {offset} is not available" @@ -133,7 +129,7 @@ def _next_halo_line(self, cells: data_alloc.NDArray) -> data_alloc.NDArray: return self._xp.setdiff1d(cell_neighbors, cells, assume_unique=True) def _find_neighbors( - self, source_indices: data_alloc.NDArray, offset: gtx.FieldOffset | str + self, source_indices: data_alloc.NDArray, offset: type[gtx.NeighborConnectivity] ) -> data_alloc.NDArray: """Get a flattened list of all (unique) neighbors to a given global index list""" assert source_indices.ndim == 1 @@ -469,7 +465,7 @@ def __call__(self, cell_to_rank: data_alloc.NDArray) -> defs.DecompositionInfo: def get_halo_constructor( process_props: defs.ProcessProperties, full_grid_size: base.HorizontalGridSize, - connectivities: dict[gtx.FieldOffset | str, data_alloc.NDArray], + connectivities: dict[type[gtx.NeighborConnectivity], data_alloc.NDArray], allocator: gtx_typing.Allocator | None, ) -> HaloConstructor: """ diff --git a/model/common/src/icon4py/model/common/grid/base.py b/model/common/src/icon4py/model/common/grid/base.py index 7282a5cba4..6da0bc5a04 100644 --- a/model/common/src/icon4py/model/common/grid/base.py +++ b/model/common/src/icon4py/model/common/grid/base.py @@ -8,7 +8,7 @@ import dataclasses import functools import logging -from collections.abc import Callable, Sequence +from collections.abc import Callable, Mapping, Sequence import gt4py.next as gtx import gt4py.next.typing as gtx_typing @@ -83,7 +83,7 @@ class Grid: UUID from icon grid files are UUID v1. """ config: GridConfig - connectivities: gtx_common.OffsetProvider + connectivities: Mapping[type[gtx.NeighborConnectivity], gtx_common.NeighborTable] start_index: Callable[[h_grid.Domain], gtx.int32] end_index: Callable[[h_grid.Domain], gtx.int32] @@ -134,10 +134,7 @@ def num_levels(self) -> int: def limited_area(self) -> bool: return self.config.limited_area - def get_connectivity(self, offset: str | gtx.FieldOffset) -> gtx_common.NeighborTable: - """Get the connectivity by its name.""" - if isinstance(offset, gtx.FieldOffset): - offset = offset.value + def get_connectivity(self, offset: type[gtx.NeighborConnectivity]) -> gtx_common.NeighborTable: if offset not in self.connectivities: raise exceptions.MissingConnectivityError( f"Missing connectivity for offset {offset} in grid {self.id}." @@ -148,15 +145,15 @@ def get_connectivity(self, offset: str | gtx.FieldOffset) -> gtx_common.Neighbor def construct_connectivity( - offset: gtx.FieldOffset, + offset: type[gtx.NeighborConnectivity], table: data_alloc.NDArray, skip_value: int | None = None, *, allocator: gtx_typing.Allocator | None = None, replace_skip_values: bool = False, ): - from_dim, dim = offset.target - to_dim = offset.source + from_dim, dim = offset.domain, gtx_common.local_dimension_of(offset) + to_dim = offset.codomain if replace_skip_values: _log.debug(f"Replacing skip values in connectivity for {dim} with max valid neighbor.") skip_value = None diff --git a/model/common/src/icon4py/model/common/grid/geometry.py b/model/common/src/icon4py/model/common/grid/geometry.py index 849588cbb7..997591afa9 100644 --- a/model/common/src/icon4py/model/common/grid/geometry.py +++ b/model/common/src/icon4py/model/common/grid/geometry.py @@ -567,7 +567,7 @@ def _register_normals_and_tangents_icosahedron(self) -> None: "y": attrs.EDGE_NORMAL_Y, "z": attrs.EDGE_NORMAL_Z, }, - connectivities={"e2c": dims.E2CDim}, + connectivities={"e2c": dims.E2C}, params={ "horizontal_start": self.grid.start_index( self._edge_domain(h_grid.Zone.LATERAL_BOUNDARY) @@ -629,7 +629,7 @@ def _register_normals_and_tangents_icosahedron(self) -> None: "y": attrs.EDGE_TANGENT_Y, "z": attrs.EDGE_TANGENT_Z, }, - connectivities={"e2c": dims.E2CDim}, + connectivities={"e2c": dims.E2C}, params={ "horizontal_start": self.grid.start_index( self._edge_domain(h_grid.Zone.LATERAL_BOUNDARY) diff --git a/model/common/src/icon4py/model/common/grid/grid_manager.py b/model/common/src/icon4py/model/common/grid/grid_manager.py index 29b09fadfc..bc7d9ffa37 100644 --- a/model/common/src/icon4py/model/common/grid/grid_manager.py +++ b/model/common/src/icon4py/model/common/grid/grid_manager.py @@ -476,13 +476,13 @@ def _construct_decomposed_grid( def _get_local_connectivities( self, - neighbor_tables_global: dict[gtx.FieldOffset, data_alloc.NDArray], - ) -> dict[gtx.FieldOffset, data_alloc.NDArray]: + neighbor_tables_global: dict[type[gtx.NeighborConnectivity], data_alloc.NDArray], + ) -> dict[type[gtx.NeighborConnectivity], data_alloc.NDArray]: if self.decomposition_info.is_distributed(): return { k: halo.global_to_local( - self._decomposition_info.global_index(k.source), - v[self._decomposition_info.global_index(k.target[0])], + self._decomposition_info.global_index(k.codomain), + v[self._decomposition_info.global_index(k.domain)], ) for k, v in neighbor_tables_global.items() } @@ -543,8 +543,8 @@ def _get_index_field( def _get_derived_connectivities( - neighbor_tables: dict[gtx.FieldOffset, data_alloc.NDArray], -) -> dict[gtx.FieldOffset, data_alloc.NDArray]: + neighbor_tables: dict[type[gtx.NeighborConnectivity], data_alloc.NDArray], +) -> dict[type[gtx.NeighborConnectivity], data_alloc.NDArray]: array_ns = data_alloc.array_namespace(next(iter(neighbor_tables.values()))) e2v_table = neighbor_tables[dims.E2V] c2v_table = neighbor_tables[dims.C2V] diff --git a/model/common/src/icon4py/model/common/grid/icon.py b/model/common/src/icon4py/model/common/grid/icon.py index 39f9c5e35a..bcf11d06e8 100644 --- a/model/common/src/icon4py/model/common/grid/icon.py +++ b/model/common/src/icon4py/model/common/grid/icon.py @@ -12,6 +12,7 @@ import gt4py.next as gtx import gt4py.next.typing as gtx_typing +from gt4py.next import common as gtx_common from icon4py.model.common import constants, dimension as dims from icon4py.model.common.grid import base, horizontal as h_grid @@ -127,14 +128,16 @@ def geometry_type(self) -> GeometryType | None: return self.grid_params.geometry_type -def _has_skip_values(offset: gtx.FieldOffset, limited_area_or_distributed: bool) -> bool: +def _has_skip_values( + offset: type[gtx.NeighborConnectivity], limited_area_or_distributed: bool +) -> bool: """ For the icosahedral global grid skip values are only present for the pentagon points. In the local area model or a distributed grid there are also skip values at the boundaries or halos when accessing neighbouring cells or edges from vertices. """ - dimension = offset.target[1] + dimension = gtx_common.local_dimension_of(offset) assert dimension.kind == gtx.DimensionKind.LOCAL, "only local dimensions can have skip values" return dimension in CONNECTIVITIES_ON_PENTAGONS or ( limited_area_or_distributed and dimension in CONNECTIVITIES_ON_BOUNDARIES @@ -142,7 +145,9 @@ def _has_skip_values(offset: gtx.FieldOffset, limited_area_or_distributed: bool) def _should_replace_skip_values( - offset: gtx.FieldOffset, keep_skip_values: bool, limited_area_or_distributed: bool + offset: type[gtx.NeighborConnectivity], + keep_skip_values: bool, + limited_area_or_distributed: bool, ) -> bool: """ Check if the skip_values in a neighbor table should be replaced. @@ -176,7 +181,7 @@ def icon_grid( id_: str, allocator: gtx_typing.Allocator | None, config: base.GridConfig, - neighbor_tables: dict[gtx.FieldOffset, data_alloc.NDArray], + neighbor_tables: dict[type[gtx.NeighborConnectivity], data_alloc.NDArray], start_index: Callable[[h_grid.Domain], gtx.int32], end_index: Callable[[h_grid.Domain], gtx.int32], grid_params: GridParams, @@ -184,7 +189,7 @@ def icon_grid( ) -> IconGrid: limited_area_or_distributed = config.limited_area or config.distributed connectivities = { - offset.value: base.construct_connectivity( + offset: base.construct_connectivity( offset, data_alloc.import_array_ns(allocator).asarray(table), skip_value=-1 if _has_skip_values(offset, limited_area_or_distributed) else None, diff --git a/model/common/src/icon4py/model/common/grid/simple.py b/model/common/src/icon4py/model/common/grid/simple.py index 4d65e53d87..1d809835a0 100644 --- a/model/common/src/icon4py/model/common/grid/simple.py +++ b/model/common/src/icon4py/model/common/grid/simple.py @@ -462,9 +462,7 @@ def simple_end_index(domain: h_grid.Domain) -> gtx.int32: } connectivities = { - offset.value: base.construct_connectivity( - offset, table, skip_value=None, allocator=allocator - ) + offset: base.construct_connectivity(offset, table, skip_value=None, allocator=allocator) for offset, table in neighbor_tables.items() } diff --git a/model/common/src/icon4py/model/common/initial_condition/analytical/linear_horizontal_tracer_advection.py b/model/common/src/icon4py/model/common/initial_condition/analytical/linear_horizontal_tracer_advection.py index 1e7b629473..ce6a2f0da0 100644 --- a/model/common/src/icon4py/model/common/initial_condition/analytical/linear_horizontal_tracer_advection.py +++ b/model/common/src/icon4py/model/common/initial_condition/analytical/linear_horizontal_tracer_advection.py @@ -14,6 +14,7 @@ import typing from typing import TYPE_CHECKING, ClassVar +from icon4py.model.common import dimension as dims from icon4py.model.common.config import config_io, options as common_conf_opt from icon4py.model.common.grid import geometry_attributes as geometry_meta, icon as icon_grid from icon4py.model.common.math import distance_array_ns @@ -367,7 +368,7 @@ def linear_horizontal_advection( vertex_y=vertex_y, cell_center_x=cell_center_x, cell_center_y=cell_center_y, - c2v_connectivity=grid.connectivities["C2V"].ndarray, + c2v_connectivity=grid.connectivities[dims.C2V].ndarray, domain_length=grid.grid_params.domain_length, domain_height=grid.grid_params.domain_height, ) @@ -419,7 +420,7 @@ def construct_reference_tracer( vertex_y=vertex_y, cell_center_x=cell_center_x, cell_center_y=cell_center_y, - c2v_connectivity=grid.connectivities["C2V"].ndarray, + c2v_connectivity=grid.connectivities[dims.C2V].ndarray, domain_length=grid.grid_params.domain_length, domain_height=grid.grid_params.domain_height, ) diff --git a/model/common/src/icon4py/model/common/interpolation/interpolation_factory.py b/model/common/src/icon4py/model/common/interpolation/interpolation_factory.py index 3a33361965..80e09ca05c 100644 --- a/model/common/src/icon4py/model/common/interpolation/interpolation_factory.py +++ b/model/common/src/icon4py/model/common/interpolation/interpolation_factory.py @@ -285,7 +285,7 @@ def _register_computed_fields(self) -> None: "dual_edge_length": geometry_attrs.DUAL_EDGE_LENGTH, "geofac_div": attrs.GEOFAC_DIV, }, - connectivities={"c2e": dims.C2EDim, "e2c": dims.E2CDim, "c2e2c": dims.C2E2CDim}, + connectivities={"c2e": dims.C2E, "e2c": dims.E2C, "c2e2c": dims.C2E2C}, params={ "horizontal_start": self._grid.start_index( cell_domain(h_grid.Zone.LATERAL_BOUNDARY_LEVEL_2) @@ -304,7 +304,7 @@ def _register_computed_fields(self) -> None: "inv_dual_edge_length": f"inverse_of_{geometry_attrs.DUAL_EDGE_LENGTH}", "owner_mask": "edge_owner_mask", }, - connectivities={"c2e": dims.C2EDim, "e2c": dims.E2CDim, "e2c2e": dims.E2C2EDim}, + connectivities={"c2e": dims.C2E, "e2c": dims.E2C, "e2c2e": dims.E2C2E}, params={ "horizontal_start": self._grid.start_index( edge_domain(h_grid.Zone.LATERAL_BOUNDARY_LEVEL_2) @@ -371,7 +371,7 @@ def _register_computed_fields(self) -> None: "cell_lon": geometry_attrs.CELL_LON, "cell_owner_mask": "cell_owner_mask", }, - connectivities={"c2e2c": dims.C2E2CDim}, + connectivities={"c2e2c": dims.C2E2C}, params={ "domain_length": self._domain_length, "domain_height": self._domain_height, @@ -403,7 +403,7 @@ def _register_computed_fields(self) -> None: "cell_areas": geometry_attrs.CELL_AREA, "cell_owner_mask": "cell_owner_mask", }, - connectivities={"c2e2c0": dims.C2E2CODim}, + connectivities={"c2e2c0": dims.C2E2CO}, params={ "horizontal_start": self.grid.start_index( cell_domain(h_grid.Zone.LATERAL_BOUNDARY_LEVEL_2) @@ -428,7 +428,7 @@ def _register_computed_fields(self) -> None: "edges_lat": geometry_attrs.EDGE_LAT, "edges_lon": geometry_attrs.EDGE_LON, }, - connectivities={"c2e": dims.C2EDim}, + connectivities={"c2e": dims.C2E}, ) self.register_provider(e_bln_c_s) @@ -447,7 +447,7 @@ def _register_computed_fields(self) -> None: "edges_lat": geometry_attrs.EDGE_LAT, "owner_mask": "edge_owner_mask", }, - connectivities={"e2c": dims.E2CDim}, + connectivities={"e2c": dims.E2C}, params={ "grid_sphere_radius": constants.EARTH_RADIUS, "horizontal_start": self.grid.start_index( @@ -470,7 +470,7 @@ def _register_computed_fields(self) -> None: "cell_areas": geometry_attrs.CELL_AREA, "cell_owner_mask": "cell_owner_mask", }, - connectivities={"c2e2c0": dims.C2E2CODim}, + connectivities={"c2e2c0": dims.C2E2CO}, params={ "horizontal_start": self.grid.start_index( cell_domain(h_grid.Zone.LATERAL_BOUNDARY_LEVEL_2) @@ -490,7 +490,7 @@ def _register_computed_fields(self) -> None: fields=(attrs.E_BLN_C_S,), domain=(dims.CellDim, dims.C2EDim), deps={}, - connectivities={"c2e": dims.C2EDim}, + connectivities={"c2e": dims.C2E}, params={}, ) self.register_provider(e_bln_c_s) @@ -502,7 +502,7 @@ def _register_computed_fields(self) -> None: deps={ "dual_edge_length": geometry_attrs.DUAL_EDGE_LENGTH, }, - connectivities={"e2c": dims.E2CDim}, + connectivities={"e2c": dims.E2C}, params={}, do_exchange=True, ) @@ -540,7 +540,7 @@ def _register_computed_fields(self) -> None: "geofac_div": attrs.GEOFAC_DIV, "c_lin_e": attrs.C_LIN_E, }, - connectivities={"c2e": dims.C2EDim, "e2c": dims.E2CDim, "c2e2c": dims.C2E2CDim}, + connectivities={"c2e": dims.C2E, "e2c": dims.E2C, "c2e2c": dims.C2E2C}, params={ "horizontal_start": self.grid.start_index( cell_domain(h_grid.Zone.LATERAL_BOUNDARY_LEVEL_2) @@ -565,10 +565,10 @@ def _register_computed_fields(self) -> None: "primal_cart_normal_z": geometry_attrs.EDGE_NORMAL_Z, }, connectivities={ - "e2c": dims.E2CDim, - "c2e": dims.C2EDim, - "c2e2c": dims.C2E2CDim, - "e2c2e": dims.E2C2EDim, + "e2c": dims.E2C, + "c2e": dims.C2E, + "c2e2c": dims.C2E2C, + "e2c2e": dims.E2C2E, }, params={ "horizontal_start_p3": self.grid.start_index( @@ -591,10 +591,10 @@ def _register_computed_fields(self) -> None: "edge_cell_length": geometry_attrs.EDGE_CELL_DISTANCE, }, connectivities={ - "v2e": dims.V2EDim, - "e2v": dims.E2VDim, - "v2c": dims.V2CDim, - "e2c": dims.E2CDim, + "v2e": dims.V2E, + "e2v": dims.E2V, + "v2c": dims.V2C, + "e2c": dims.E2C, }, params={ "horizontal_start": self.grid.start_index( @@ -623,7 +623,7 @@ def _register_computed_fields(self) -> None: "edge_normal_z": geometry_attrs.EDGE_NORMAL_Z, "scale_factor": attrs.RBF_SCALE_CELL, }, - connectivities={"rbf_offset": dims.C2E2C2EDim}, + connectivities={"rbf_offset": dims.C2E2C2E}, params={ "rbf_kernel": self._config.rbf_kernel_cell.value, "geometry_type": self._grid.grid_params.geometry_type.value, @@ -659,7 +659,7 @@ def _register_computed_fields(self) -> None: "edge_dual_normal_v": geometry_attrs.EDGE_DUAL_V, "scale_factor": attrs.RBF_SCALE_EDGE, }, - connectivities={"rbf_offset": dims.E2C2EDim}, + connectivities={"rbf_offset": dims.E2C2E}, params={ "rbf_kernel": self._config.rbf_kernel_edge.value, "geometry_type": self._grid.grid_params.geometry_type.value, @@ -696,7 +696,7 @@ def _register_computed_fields(self) -> None: "edge_normal_z": geometry_attrs.EDGE_NORMAL_Z, "scale_factor": attrs.RBF_SCALE_VERTEX, }, - connectivities={"rbf_offset": dims.V2EDim}, + connectivities={"rbf_offset": dims.V2E}, params={ "rbf_kernel": self._config.rbf_kernel_vertex.value, "geometry_type": self._grid.grid_params.geometry_type.value, diff --git a/model/common/src/icon4py/model/common/metrics/compute_zdiff_gradp.py b/model/common/src/icon4py/model/common/metrics/compute_zdiff_gradp.py index ecca4ac982..4db2768610 100644 --- a/model/common/src/icon4py/model/common/metrics/compute_zdiff_gradp.py +++ b/model/common/src/icon4py/model/common/metrics/compute_zdiff_gradp.py @@ -55,7 +55,7 @@ def compute_zdiff_gradp( # noqa: PLR0912 [too-many-branches] >>> z_ifc_off=z_ifc_off, >>> domain={EdgeDim: (horizontal_start, nedges), KDim: (0, nlev)}, >>> out=z_ifc_off_koff, - >>> offset_provider={"Koff": icon_grid.get_offset_provider("Koff")} + >>> offset_provider={} >>> ) """ diff --git a/model/common/src/icon4py/model/common/metrics/metrics_factory.py b/model/common/src/icon4py/model/common/metrics/metrics_factory.py index 076b86864c..e7fbc5affa 100644 --- a/model/common/src/icon4py/model/common/metrics/metrics_factory.py +++ b/model/common/src/icon4py/model/common/metrics/metrics_factory.py @@ -230,7 +230,7 @@ def _register_computed_fields(self) -> None: # noqa: PLR0915 [too-many-statemen "cell_areas": geometry_attrs.CELL_AREA, "geofac_n2s": interpolation_attributes.GEOFAC_N2S, }, - connectivities={"c2e2co": dims.C2E2CODim}, + connectivities={"c2e2co": dims.C2E2CO}, params={ "nflatlev": self._vertical_grid.nflatlev, "model_top_height": self._vertical_grid.config.model_top_height, @@ -616,7 +616,7 @@ def _register_computed_fields(self) -> None: # noqa: PLR0915 [too-many-statemen compute_exner_w_implicit_weight_parameter_np = factory.NumpyDataProvider( func=mf.compute_exner_w_implicit_weight_parameter, domain=(dims.CellDim,), - connectivities={"c2e": dims.C2EDim}, + connectivities={"c2e": dims.C2E}, fields=(attrs.EXNER_W_IMPLICIT_WEIGHT_PARAMETER,), deps={ "vct_a": "vct_a", @@ -728,7 +728,7 @@ def _register_computed_fields(self) -> None: # noqa: PLR0915 [too-many-statemen "z_ifc": attrs.CELL_HEIGHT_ON_HALF_LEVEL, "k_lev": "k_lev", }, - connectivities={"e2c": dims.E2CDim}, + connectivities={"e2c": dims.E2C}, domain={ dims.EdgeDim: ( edge_domain(h_grid.Zone.LATERAL_BOUNDARY_LEVEL_2), @@ -841,7 +841,7 @@ def _register_computed_fields(self) -> None: # noqa: PLR0915 [too-many-statemen "flat_idx": attrs.FLAT_IDX_MAX, "topography": "topography", }, - connectivities={"e2c": dims.E2CDim}, + connectivities={"e2c": dims.E2C}, domain=(dims.EdgeDim, dims.E2CDim, dims.KDim), fields=( attrs.ZDIFF_GRADP, @@ -1040,7 +1040,7 @@ def _register_computed_fields(self) -> None: # noqa: PLR0915 [too-many-statemen deps={ "z_mc": attrs.Z_MC, }, - connectivities={"c2e2c": dims.C2E2CDim}, + connectivities={"c2e2c": dims.C2E2C}, domain=(dims.CellDim,), fields=(attrs.MAX_NBHGT,), params={ @@ -1059,7 +1059,7 @@ def _register_computed_fields(self) -> None: # noqa: PLR0915 [too-many-statemen "maxslp_avg": attrs.MAXSLP_AVG, "maxhgtd_avg": attrs.MAXHGTD_AVG, }, - connectivities={"c2e2c": dims.C2E2CDim}, + connectivities={"c2e2c": dims.C2E2C}, domain=(dims.CellDim, dims.KDim), fields=(attrs.ZD_DIFFCOEF,), params={ @@ -1083,7 +1083,7 @@ def _register_computed_fields(self) -> None: # noqa: PLR0915 [too-many-statemen "maxslp_avg": attrs.MAXSLP_AVG, "maxhgtd_avg": attrs.MAXHGTD_AVG, }, - connectivities={"c2e2c": dims.C2E2CDim}, + connectivities={"c2e2c": dims.C2E2C}, domain=(dims.CellDim, dims.C2E2CDim, dims.KDim), fields=( attrs.ZD_INTCOEF, diff --git a/model/common/src/icon4py/model/common/states/factory.py b/model/common/src/icon4py/model/common/states/factory.py index 4cc7d22053..77cfbff4ca 100644 --- a/model/common/src/icon4py/model/common/states/factory.py +++ b/model/common/src/icon4py/model/common/states/factory.py @@ -439,7 +439,9 @@ def _unravel_output_fields(self): return out_fields # TODO(): do we need that here? - def _get_offset_providers(self, grid: icon_grid.IconGrid) -> dict[str, gtx.FieldOffset]: + def _get_offset_providers( + self, grid: icon_grid.IconGrid + ) -> dict[type[gtx.NeighborConnectivity], gtx_common.NeighborTable]: offset_providers = {} for dim in self._dims: if dim.kind == gtx.DimensionKind.HORIZONTAL: @@ -454,7 +456,8 @@ def _get_offset_providers(self, grid: icon_grid.IconGrid) -> dict[str, gtx.Field vertical_offsets = { k: v for k, v in grid.connectivities.items() - if isinstance(v, gtx_common.DimensionMeta) and v.kind == gtx.DimensionKind.VERTICAL + if isinstance(v, gtx_common.DimensionMeta) + and v.kind == gtx.DimensionKind.VERTICAL } offset_providers.update(vertical_offsets) # used for different compute backend in function call @@ -530,7 +533,9 @@ def _allocate( # TODO(halungge): this can be simplified when completely disentangling vertical and horizontal grid. # the IconGrid should then only contain horizontal connectivities and no longer any Koff which should be moved to the VerticalGrid - def _get_offset_providers(self, grid: icon_grid.IconGrid) -> dict[str, gtx.FieldOffset]: + def _get_offset_providers( + self, grid: icon_grid.IconGrid + ) -> dict[type[gtx.NeighborConnectivity], gtx_common.NeighborTable]: offset_providers = {} for dim in self._domain: if dim.kind == gtx.DimensionKind.HORIZONTAL: @@ -546,7 +551,8 @@ def _get_offset_providers(self, grid: icon_grid.IconGrid) -> dict[str, gtx.Field vertical_offsets = { k: v for k, v in grid.connectivities.items() - if isinstance(v, gtx_common.DimensionMeta) and v.kind == gtx.DimensionKind.VERTICAL + if isinstance(v, gtx_common.DimensionMeta) + and v.kind == gtx.DimensionKind.VERTICAL } offset_providers.update(vertical_offsets) return offset_providers @@ -630,8 +636,8 @@ class NumpyDataProvider(FieldProvider, NeedsExchange): fields: Seq[str] names under which the results fo the function will be registered deps: dict[str, str] input fields used for computing this stencil: the key is the variable name used in the function and the value the name of the field it depends on. - connectivities: dict[str, Dimension] dict where the key is the variable named used in the - function and the value the sparse Dimension of the connectivity field + connectivities: dict[str, type[NeighborConnectivity]] dict where the key is the variable + name used in the function and the value the connectivity whose table is passed params: scalar arguments for the function do_exchange: a flag that governs whether or not a halo exchange is needed after the field has been computed. Defaults to False """ @@ -643,7 +649,7 @@ def __init__( domain: dict[gtx.Dimension, tuple[DomainType, DomainType]] | tuple[gtx.Dimension, ...], fields: Sequence[str], deps: dict[str, str], - connectivities: dict[str, gtx.Dimension] | None = None, + connectivities: dict[str, type[gtx.NeighborConnectivity]] | None = None, params: dict[str, state_utils.ScalarType] | None = None, do_exchange: bool = False, ): @@ -688,7 +694,7 @@ def _compute( for k, v in self._dependencies.items() } offsets = { - k: grid_provider.grid.get_connectivity(v.value).ndarray + k: grid_provider.grid.get_connectivity(v).ndarray for k, v in self._connectivities.items() } args.update(offsets) diff --git a/model/common/tests/common/decomposition/mpi_tests/test_mpi_decomposition.py b/model/common/tests/common/decomposition/mpi_tests/test_mpi_decomposition.py index 6042d7dba2..9cdcb0a481 100644 --- a/model/common/tests/common/decomposition/mpi_tests/test_mpi_decomposition.py +++ b/model/common/tests/common/decomposition/mpi_tests/test_mpi_decomposition.py @@ -335,7 +335,7 @@ def test_halo_exchange_for_sparse_field( # noqa: PLR0917 [too-many-positional-a edge_orientation, area, out=result, - offset_provider={"C2E": icon_grid.get_connectivity("C2E")}, + offset_provider={dims.C2E: icon_grid.get_connectivity(dims.C2E)}, ) _log.info( f"{process_props.rank}/{process_props.comm_size}: size of computed field {result.asnumpy().shape}" diff --git a/model/common/tests/common/decomposition/unit_tests/test_halo.py b/model/common/tests/common/decomposition/unit_tests/test_halo.py index d7f50edafa..771e46ef14 100644 --- a/model/common/tests/common/decomposition/unit_tests/test_halo.py +++ b/model/common/tests/common/decomposition/unit_tests/test_halo.py @@ -218,10 +218,10 @@ def test_global_to_local_index(offset, rank): process_props = dummy_four_ranks(rank) halo_constructor = halo.IconLikeHaloConstructor(process_props, neighbor_tables) decomposition_info = halo_constructor(utils.SIMPLE_DISTRIBUTION) - source_indices_on_local_grid = decomposition_info.global_index(offset.target[0]) + source_indices_on_local_grid = decomposition_info.global_index(offset.domain) - offset_full_grid = grid.connectivities[offset.value].ndarray[source_indices_on_local_grid] - neighbor_dim = offset.source + offset_full_grid = grid.connectivities[offset].ndarray[source_indices_on_local_grid] + neighbor_dim = offset.codomain neighbor_index_full_grid = decomposition_info.global_index(neighbor_dim) local_offset = halo.global_to_local( diff --git a/model/common/tests/common/grid/unit_tests/test_grid_manager.py b/model/common/tests/common/grid/unit_tests/test_grid_manager.py index c3741ddb17..8e7b92c053 100644 --- a/model/common/tests/common/grid/unit_tests/test_grid_manager.py +++ b/model/common/tests/common/grid/unit_tests/test_grid_manager.py @@ -15,6 +15,7 @@ import gt4py.next.typing as gtx_typing import numpy as np import pytest +from gt4py.next import common as gtx_common import icon4py.model.common.grid.gridfile from icon4py.model.common import dimension as dims, model_backends @@ -83,7 +84,7 @@ def test_grid_manager_eval_v2e( # 6 neighbors hence there are "Missing values" in the grid file # they get substituted by the "last valid index" in preprocessing step in icon. assert not has_invalid_index(seralized_v2e) - v2e_table = grid.get_connectivity("V2E").asnumpy() + v2e_table = grid.get_connectivity(dims.V2E).asnumpy() # Torus grids have no pentagon points and no boundaries hence no invalid # indexes (while REGIONAL and GLOBAL grids can have) assert ( @@ -124,7 +125,7 @@ def test_grid_manager_eval_v2c( ) -> None: grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid serialized_v2c = data_alloc.as_numpy(grid_savepoint.v2c()) - v2c_table = grid.get_connectivity("V2C").asnumpy() + v2c_table = grid.get_connectivity(dims.V2C).asnumpy() # there are vertices that have less than 6 neighboring cells: either pentagon points or # vertices at the boundary of the domain for a limited area mode # hence in the grid file there are "missing values" @@ -179,7 +180,7 @@ def test_grid_manager_eval_e2v( grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid serialized_e2v = data_alloc.as_numpy(grid_savepoint.e2v()) - e2v_table = grid.get_connectivity("E2V").asnumpy() + e2v_table = grid.get_connectivity(dims.E2V).asnumpy() # all vertices in the system have to neighboring edges, there no edges that point nowhere # hence this connectivity has no "missing values" in the grid file assert not has_invalid_index(serialized_e2v) @@ -202,7 +203,7 @@ def test_grid_manager_eval_e2c( grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid serialized_e2c = data_alloc.as_numpy(grid_savepoint.e2c()) - e2c_table = grid.get_connectivity("E2C").asnumpy() + e2c_table = grid.get_connectivity(dims.E2C).asnumpy() assert has_invalid_index(serialized_e2c) == grid.limited_area assert has_invalid_index(e2c_table) == grid.limited_area assert np.allclose(e2c_table, serialized_e2c) @@ -219,7 +220,7 @@ def test_grid_manager_eval_c2e( grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid serialized_c2e = data_alloc.as_numpy(grid_savepoint.c2e()) - c2e_table = grid.get_connectivity("C2E").asnumpy() + c2e_table = grid.get_connectivity(dims.C2E).asnumpy() # no cells with less than 3 neighboring edges exist, otherwise the cell is not there in the # first place # hence there are no "missing values" in the grid file @@ -238,7 +239,7 @@ def test_grid_manager_eval_c2e2c( ) -> None: grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid assert np.allclose( - grid.get_connectivity("C2E2C").asnumpy(), + grid.get_connectivity(dims.C2E2C).asnumpy(), data_alloc.as_numpy(grid_savepoint.c2e2c()), ) @@ -253,8 +254,8 @@ def test_grid_manager_eval_c2e2cO( grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid serialized_grid = grid_savepoint.construct_icon_grid(backend=backend) assert np.allclose( - grid.get_connectivity("C2E2CO").asnumpy(), - serialized_grid.get_connectivity("C2E2CO").asnumpy(), + grid.get_connectivity(dims.C2E2CO).asnumpy(), + serialized_grid.get_connectivity(dims.C2E2CO).asnumpy(), ) @@ -268,12 +269,12 @@ def test_grid_manager_eval_e2c2e( ) -> None: grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid serialized_grid = grid_savepoint.construct_icon_grid(backend=backend) - serialized_e2c2e = serialized_grid.get_connectivity("E2C2E").asnumpy() - serialized_e2c2eO = serialized_grid.get_connectivity("E2C2EO").asnumpy() + serialized_e2c2e = serialized_grid.get_connectivity(dims.E2C2E).asnumpy() + serialized_e2c2eO = serialized_grid.get_connectivity(dims.E2C2EO).asnumpy() assert has_invalid_index(serialized_e2c2e) == grid.limited_area - e2c2e_table = grid.get_connectivity("E2C2E").asnumpy() - e2c2eO_table = grid.get_connectivity("E2C2EO").asnumpy() + e2c2e_table = grid.get_connectivity(dims.E2C2E).asnumpy() + e2c2eO_table = grid.get_connectivity(dims.E2C2EO).asnumpy() assert has_invalid_index(e2c2e_table) == grid.limited_area # ICON calculates diamond edges only from rl_start = 2 (lateral_boundary(dims.EdgeDim) + 1 for # boundaries all values are INVALID even though the half diamond exists (see mo_model_domimp_setup.f90 ll 163ff.) @@ -295,13 +296,13 @@ def test_grid_manager_eval_e2c2v( serialized_ref = data_alloc.as_numpy(grid_savepoint.e2c2v()) # the "far" (adjacent to edge normal ) is not always there, because ICON only calculates those starting from # (lateral_boundary(dims.EdgeDim) + 1) to end(dims.EdgeDim) (see mo_intp_coeffs.f90) and only for owned cells - table = grid.get_connectivity("E2C2V").asnumpy() + table = grid.get_connectivity(dims.E2C2V).asnumpy() start_index = grid.start_index( h_grid.domain(dims.EdgeDim)(h_grid.Zone.LATERAL_BOUNDARY_LEVEL_2) ) # e2c2e in ICON (quad_idx) has a different neighbor ordering than the e2c2e constructed in grid_manager.py assert_up_to_order(table, serialized_ref, start_index) - assert np.allclose(table[:, :2], grid.get_connectivity("E2V").asnumpy()) + assert np.allclose(table[:, :2], grid.get_connectivity(dims.E2V).asnumpy()) @pytest.mark.datatest @@ -312,7 +313,7 @@ def test_grid_manager_eval_c2v( backend: gtx_typing.Backend, ) -> None: grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid - c2v = grid.get_connectivity("C2V").asnumpy() + c2v = grid.get_connectivity(dims.C2V).asnumpy() assert np.allclose(c2v, data_alloc.as_numpy(grid_savepoint.c2v())) @@ -406,10 +407,10 @@ def test_grid_manager_eval_c2e2c2e( grid = utils.run_grid_manager(experiment.grid, keep_skip_values=True, backend=backend).grid serialized_grid = grid_savepoint.construct_icon_grid(backend=backend) assert np.allclose( - grid.get_connectivity("C2E2C2E").asnumpy(), - serialized_grid.get_connectivity("C2E2C2E").asnumpy(), + grid.get_connectivity(dims.C2E2C2E).asnumpy(), + serialized_grid.get_connectivity(dims.C2E2C2E).asnumpy(), ) - assert grid.get_connectivity("C2E2C2E").asnumpy().shape == (grid.num_cells, 9) + assert grid.get_connectivity(dims.C2E2C2E).asnumpy().shape == (grid.num_cells, 9) # TODO (halungge): check EXCOAIM APE with new serialized data ( standard grid, start_idx/end_idx arrays @@ -609,12 +610,12 @@ def test_decomposition_info_single_rank( dims.E2C2E, dims.E2C2EO, ], - ids=lambda offset: offset.value, + ids=lambda offset: offset.__name__, ) def test_local_connectivity( rank: int, caplog: Iterator, - field_offset: gtx.FieldOffset, + field_offset: type[gtx.NeighborConnectivity], backend_like: model_backends.BackendLike, ) -> None: process_props = decomp_utils.DummyProps(rank=rank) @@ -637,22 +638,22 @@ def test_local_connectivity( assert ( connectivity.shape[0] == decomposition_info.global_index( - field_offset.target[0], decomp_defs.DecompositionInfo.EntryType.ALL + field_offset.domain, decomp_defs.DecompositionInfo.EntryType.ALL ).size ), "connectivity shapes do not match" # all neighbor indices are valid local indices max_local_index = np.max( decomposition_info.local_index( - field_offset.source, decomp_defs.DecompositionInfo.EntryType.ALL + field_offset.codomain, decomp_defs.DecompositionInfo.EntryType.ALL ) ) assert np.max(connectivity) == max_local_index, ( f"max value in the connectivity is {np.max(connectivity)} is larger than the local patch size {max_local_index}" ) # - outer halo entries have SKIP_VALUE neighbors (depends on offsets) - neighbor_dim = field_offset.target[1] # type: ignore [misc] - dim = field_offset.target[0] + neighbor_dim = gtx_common.local_dimension_of(field_offset) + dim = field_offset.domain last_halo_level = ( decomp_defs.DecompositionFlag.THIRD_HALO_LEVEL if dim == dims.EdgeDim diff --git a/model/common/tests/common/grid/unit_tests/test_icon.py b/model/common/tests/common/grid/unit_tests/test_icon.py index 816b395209..20af08a8a3 100644 --- a/model/common/tests/common/grid/unit_tests/test_icon.py +++ b/model/common/tests/common/grid/unit_tests/test_icon.py @@ -185,10 +185,10 @@ def test_grid_size(icon_grid: base_grid.Grid) -> None: "grid_description", (test_defs.Grids.MCH_CH_R04B09_DSL, test_defs.Grids.R02B04_GLOBAL), ) -@pytest.mark.parametrize("offset", (utils.horizontal_offsets()), ids=lambda x: x.value) +@pytest.mark.parametrize("offset", (utils.horizontal_offsets()), ids=lambda x: x.__name__) def test_when_keep_skip_value_then_neighbor_table_matches_config( grid_description: test_defs.GridDescription, - offset: gtx.FieldOffset, + offset: type[gtx.NeighborConnectivity], backend: gtx_typing.Backend, ) -> None: grid = utils.run_grid_manager(grid_description, keep_skip_values=True, backend=backend).grid diff --git a/model/common/tests/common/grid/unit_tests/test_topography.py b/model/common/tests/common/grid/unit_tests/test_topography.py index 383a1b8d69..c94bb23571 100644 --- a/model/common/tests/common/grid/unit_tests/test_topography.py +++ b/model/common/tests/common/grid/unit_tests/test_topography.py @@ -11,7 +11,7 @@ import pytest -from icon4py.model.common import topography as topo +from icon4py.model.common import dimension as dims, topography as topo from icon4py.model.common.decomposition import definitions as decomposition from icon4py.model.testing import test_utils from icon4py.model.testing.fixtures import * # noqa: F403 @@ -47,7 +47,7 @@ def test_topography_smoothing_with_serialized_data( topography=topography.ndarray, cell_areas=cell_geometry.area.ndarray, geofac_n2s=geofac_n2s.ndarray, - c2e2co=icon_grid.get_connectivity("C2E2CO").ndarray, + c2e2co=icon_grid.get_connectivity(dims.C2E2CO).ndarray, num_iterations=num_iterations, exchange=decomposition.SingleNodeExchange(), ) diff --git a/model/common/tests/common/grid/unit_tests/test_vertical.py b/model/common/tests/common/grid/unit_tests/test_vertical.py index c44fac220e..2f3c698a20 100644 --- a/model/common/tests/common/grid/unit_tests/test_vertical.py +++ b/model/common/tests/common/grid/unit_tests/test_vertical.py @@ -384,7 +384,7 @@ def test_compute_vertical_coordinate( # noqa: PLR0917 [too-many-positional-argu topography=topography.ndarray, cell_areas=cell_geometry.area.ndarray, geofac_n2s=geofac_n2s.ndarray, - c2e2co=icon_grid.get_connectivity("C2E2CO").ndarray, + c2e2co=icon_grid.get_connectivity(dims.C2E2CO).ndarray, nflatlev=vertical_geometry.nflatlev, model_top_height=vertical_config.model_top_height, SLEVE_decay_scale_1=vertical_config.SLEVE_decay_scale_1, diff --git a/model/common/tests/common/grid/utils.py b/model/common/tests/common/grid/utils.py index 9c21919c09..ad153d0699 100644 --- a/model/common/tests/common/grid/utils.py +++ b/model/common/tests/common/grid/utils.py @@ -21,9 +21,9 @@ managers: dict[str, gm.GridManager] = {} -def horizontal_offsets() -> Iterator[gtx.FieldOffset]: +def horizontal_offsets() -> Iterator[type[gtx.NeighborConnectivity]]: for d in vars(dims).values(): - if isinstance(d, gtx.FieldOffset) and len(d.target) == 2: + if isinstance(d, type) and issubclass(d, gtx.NeighborConnectivity): yield d diff --git a/model/common/tests/common/interpolation/integration_tests/test_edge_2_cell_vector_rbf_interpolation.py b/model/common/tests/common/interpolation/integration_tests/test_edge_2_cell_vector_rbf_interpolation.py index 52581eddd8..57dfd05c5c 100644 --- a/model/common/tests/common/interpolation/integration_tests/test_edge_2_cell_vector_rbf_interpolation.py +++ b/model/common/tests/common/interpolation/integration_tests/test_edge_2_cell_vector_rbf_interpolation.py @@ -84,7 +84,7 @@ def test_edge_2_cell_vector_rbf_interpolation( vertical_start=0, vertical_end=icon_grid.num_levels, offset_provider={ - "C2E2C2E": icon_grid.get_connectivity("C2E2C2E"), + dims.C2E2C2E: icon_grid.get_connectivity(dims.C2E2C2E), }, ) diff --git a/model/common/tests/common/interpolation/unit_tests/test_interpolation_fields.py b/model/common/tests/common/interpolation/unit_tests/test_interpolation_fields.py index 3d8478159f..3282714692 100644 --- a/model/common/tests/common/interpolation/unit_tests/test_interpolation_fields.py +++ b/model/common/tests/common/interpolation/unit_tests/test_interpolation_fields.py @@ -106,7 +106,7 @@ def test_compute_geofac_div( edge_orientation=edge_orientation, area=area, out=geofac_div, - offset_provider={"C2E": mesh.get_connectivity("C2E")}, + offset_provider={dims.C2E: mesh.get_connectivity(dims.C2E)}, ) assert test_helpers.dallclose(geofac_div.asnumpy(), geofac_div_ref.asnumpy()) @@ -138,7 +138,7 @@ def test_compute_geofac_rot( owner_mask, out=geofac_rot, domain={dims.VertexDim: (horizontal_start, horizontal_end)}, - offset_provider={"V2E": mesh.get_connectivity("V2E")}, + offset_provider={dims.V2E: mesh.get_connectivity(dims.V2E)}, ) assert test_helpers.dallclose(geofac_rot.asnumpy(), geofac_rot_ref.asnumpy()) diff --git a/model/common/tests/common/math/stencil_tests/test_cell_horizontal_gradients_by_green_gauss_method.py b/model/common/tests/common/math/stencil_tests/test_cell_horizontal_gradients_by_green_gauss_method.py index a82d6438ff..8886c1579c 100644 --- a/model/common/tests/common/math/stencil_tests/test_cell_horizontal_gradients_by_green_gauss_method.py +++ b/model/common/tests/common/math/stencil_tests/test_cell_horizontal_gradients_by_green_gauss_method.py @@ -23,7 +23,7 @@ def cell_horizontal_gradients_by_green_gauss_method_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], scalar_field: np.ndarray, geofac_grg_x: np.ndarray, geofac_grg_y: np.ndarray, diff --git a/model/common/tests/common/metrics/unit_tests/test_compute_diffusion_metrics.py b/model/common/tests/common/metrics/unit_tests/test_compute_diffusion_metrics.py index e8f826b73d..a7ac233ad2 100644 --- a/model/common/tests/common/metrics/unit_tests/test_compute_diffusion_metrics.py +++ b/model/common/tests/common/metrics/unit_tests/test_compute_diffusion_metrics.py @@ -85,7 +85,7 @@ def test_compute_diffusion_mask_and_coeff( # noqa: PLR0917 [too-many-positional horizontal_end=icon_grid.num_cells, vertical_start=0, vertical_end=nlev, - offset_provider={"C2E": icon_grid.get_connectivity("C2E")}, + offset_provider={dims.C2E: icon_grid.get_connectivity(dims.C2E)}, ) compute_weighted_cell_neighbor_sum.with_backend(backend)( @@ -99,7 +99,7 @@ def test_compute_diffusion_mask_and_coeff( # noqa: PLR0917 [too-many-positional vertical_start=0, vertical_end=nlev, offset_provider={ - "C2E2CO": icon_grid.get_connectivity("C2E2CO"), + dims.C2E2CO: icon_grid.get_connectivity(dims.C2E2CO), }, ) @@ -108,7 +108,7 @@ def test_compute_diffusion_mask_and_coeff( # noqa: PLR0917 [too-many-positional max_nbhgt=max_nbhgt, horizontal_start=cell_nudging, horizontal_end=icon_grid.num_cells, - offset_provider={"C2E2C": icon_grid.get_connectivity("C2E2C")}, + offset_provider={dims.C2E2C: icon_grid.get_connectivity(dims.C2E2C)}, ) zd_diffcoef = compute_diffusion_mask_and_coef( @@ -168,7 +168,7 @@ def test_compute_diffusion_intcoef_and_vertoffset( # noqa: PLR0917 [too-many-po horizontal_end=icon_grid.num_cells, vertical_start=0, vertical_end=nlev, - offset_provider={"C2E": icon_grid.get_connectivity("C2E")}, + offset_provider={dims.C2E: icon_grid.get_connectivity(dims.C2E)}, ) compute_weighted_cell_neighbor_sum.with_backend(backend)( @@ -182,7 +182,7 @@ def test_compute_diffusion_intcoef_and_vertoffset( # noqa: PLR0917 [too-many-po vertical_start=0, vertical_end=nlev, offset_provider={ - "C2E2CO": icon_grid.get_connectivity("C2E2CO"), + dims.C2E2CO: icon_grid.get_connectivity(dims.C2E2CO), }, ) @@ -191,7 +191,7 @@ def test_compute_diffusion_intcoef_and_vertoffset( # noqa: PLR0917 [too-many-po max_nbhgt=max_nbhgt, horizontal_start=cell_nudging, horizontal_end=icon_grid.num_cells, - offset_provider={"C2E2C": icon_grid.get_connectivity("C2E2C")}, + offset_provider={dims.C2E2C: icon_grid.get_connectivity(dims.C2E2C)}, ) zd_intcoef, zd_vertoffset = compute_diffusion_intcoef_and_vertoffset( diff --git a/model/common/tests/common/metrics/unit_tests/test_compute_zdiff_gradp.py b/model/common/tests/common/metrics/unit_tests/test_compute_zdiff_gradp.py index 91b2a29b9f..9ec7346995 100644 --- a/model/common/tests/common/metrics/unit_tests/test_compute_zdiff_gradp.py +++ b/model/common/tests/common/metrics/unit_tests/test_compute_zdiff_gradp.py @@ -63,7 +63,7 @@ def test_compute_zdiff_gradp( start_nudging = icon_grid.start_index(edge_domain(h_grid.Zone.NUDGING_LEVEL_2)) flat_idx_np = compute_flat_max_idx( - e2c=icon_grid.get_connectivity("E2C").ndarray, + e2c=icon_grid.get_connectivity(dims.E2C).ndarray, z_mc=z_mc.ndarray, c_lin_e=c_lin_e.ndarray, z_ifc=z_ifc.ndarray, @@ -72,7 +72,7 @@ def test_compute_zdiff_gradp( ) zdiff_gradp_full_field, vertoffset_gradp_full_field = compute_zdiff_gradp( - e2c=icon_grid.get_connectivity("E2C").ndarray, + e2c=icon_grid.get_connectivity(dims.E2C).ndarray, z_mc=z_mc.ndarray, c_lin_e=c_lin_e.ndarray, z_ifc=metrics_savepoint.z_ifc().ndarray, diff --git a/model/common/tests/common/metrics/unit_tests/test_metric_fields.py b/model/common/tests/common/metrics/unit_tests/test_metric_fields.py index 6c224a70d8..43e00fa38f 100644 --- a/model/common/tests/common/metrics/unit_tests/test_metric_fields.py +++ b/model/common/tests/common/metrics/unit_tests/test_metric_fields.py @@ -211,7 +211,7 @@ def test_compute_exner_w_explicit_weight_parameter( exner_w_explicit_weight_parameter=exner_w_explicit_weight_parameter_full, horizontal_start=0, horizontal_end=icon_grid.num_cells, - offset_provider={"C2E": icon_grid.get_connectivity("C2E")}, + offset_provider={dims.C2E: icon_grid.get_connectivity(dims.C2E)}, ) assert testing_helpers.dallclose( @@ -238,7 +238,7 @@ def test_compute_exner_exfac( metrics_savepoint.ddxn_z_full(), grid_savepoint.dual_edge_length(), out=(max_slp, max_hgtd), - offset_provider={"C2E": icon_grid.get_connectivity("C2E")}, + offset_provider={dims.C2E: icon_grid.get_connectivity(dims.C2E)}, domain={ dims.CellDim: (horizontal_start, icon_grid.num_cells), dims.KDim: (0, icon_grid.num_levels), @@ -297,7 +297,7 @@ def test_compute_exner_w_implicit_weight_parameter( # noqa: PLR0917 [too-many-p horizontal_end=horizontal_end, vertical_start=vertical_start, vertical_end=vertical_end, - offset_provider={"E2C": icon_grid.get_connectivity("E2C")}, + offset_provider={dims.E2C: icon_grid.get_connectivity(dims.E2C)}, ) horizontal_start_edge = icon_grid.start_index( @@ -316,8 +316,8 @@ def test_compute_exner_w_implicit_weight_parameter( # noqa: PLR0917 [too-many-p vertical_start=vertical_start, vertical_end=vertical_end, offset_provider={ - "E2V": icon_grid.get_connectivity("E2V"), - "V2C": icon_grid.get_connectivity("V2C"), + dims.E2V: icon_grid.get_connectivity(dims.E2V), + dims.V2C: icon_grid.get_connectivity(dims.V2C), }, ) @@ -361,7 +361,7 @@ def test_compute_wgtfac_e( horizontal_end=icon_grid.num_edges, vertical_start=0, vertical_end=icon_grid.num_levels + 1, - offset_provider={"E2C": icon_grid.get_connectivity("E2C")}, + offset_provider={dims.E2C: icon_grid.get_connectivity(dims.E2C)}, ) assert testing_helpers.dallclose(wgtfac_e.asnumpy(), wgtfac_e_ref.asnumpy()) @@ -393,7 +393,7 @@ def test_compute_pressure_gradient_downward_extrapolation_mask_distance( start_edge_nudging_2 = icon_grid.start_index(edge_domain(horizontal.Zone.NUDGING_LEVEL_2)) flat_idx_max = mf.compute_flat_max_idx( - e2c=icon_grid.get_connectivity("E2C").ndarray, + e2c=icon_grid.get_connectivity(dims.E2C).ndarray, z_mc=z_mc.ndarray, c_lin_e=c_lin_e.ndarray, z_ifc=z_ifc.ndarray, @@ -418,7 +418,7 @@ def test_compute_pressure_gradient_downward_extrapolation_mask_distance( vertical_start=0, vertical_end=icon_grid.num_levels, offset_provider={ - "E2C": icon_grid.get_connectivity("E2C"), + dims.E2C: icon_grid.get_connectivity(dims.E2C), }, ) diff --git a/model/common/tests/common/metrics/unit_tests/test_reference_atmosphere.py b/model/common/tests/common/metrics/unit_tests/test_reference_atmosphere.py index a2869adafd..73f24d5b2f 100644 --- a/model/common/tests/common/metrics/unit_tests/test_reference_atmosphere.py +++ b/model/common/tests/common/metrics/unit_tests/test_reference_atmosphere.py @@ -166,7 +166,7 @@ def test_compute_reference_atmosphere_on_full_level_edge_fields( horizontal_end=(gtx.int32(icon_grid.num_edges)), vertical_start=(gtx.int32(0)), vertical_end=(gtx.int32(icon_grid.num_levels)), - offset_provider={"E2C": icon_grid.get_connectivity("E2C")}, + offset_provider={dims.E2C: icon_grid.get_connectivity(dims.E2C)}, ) assert stencil_tests.dallclose(rho_ref_me.asnumpy(), rho_ref_me_ref.asnumpy(), rtol=1e-10) assert stencil_tests.dallclose(theta_ref_me.asnumpy(), theta_ref_me_ref.asnumpy()) diff --git a/model/driver/src/icon4py/model/driver/driver_io.py b/model/driver/src/icon4py/model/driver/driver_io.py index a101bd5032..904a12d343 100644 --- a/model/driver/src/icon4py/model/driver/driver_io.py +++ b/model/driver/src/icon4py/model/driver/driver_io.py @@ -203,7 +203,7 @@ def compute( horizontal_end=end_cell_end, vertical_start=0, vertical_end=num_levels, - offset_provider={"C2E2C2E": self._grid.get_connectivity("C2E2C2E")}, + offset_provider={dims.C2E2C2E: self._grid.get_connectivity(dims.C2E2C2E)}, ) compute_pressure.compute_surface_and_hydrostatic_pressure.with_backend(backend)( diff --git a/model/testing/src/icon4py/model/testing/reference_funcs.py b/model/testing/src/icon4py/model/testing/reference_funcs.py index 30d04f5f64..7f98a71874 100644 --- a/model/testing/src/icon4py/model/testing/reference_funcs.py +++ b/model/testing/src/icon4py/model/testing/reference_funcs.py @@ -34,7 +34,9 @@ def enhanced_smagorinski_factor_numpy( def nabla2_on_cell_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], psi_c: np.ndarray, geofac_n2s: np.ndarray + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], + psi_c: np.ndarray, + geofac_n2s: np.ndarray, ) -> np.ndarray: c2e2cO = connectivities[dims.C2E2CO] nabla2_psi_c = np.sum(np.where((c2e2cO != -1), psi_c[c2e2cO] * geofac_n2s, 0), axis=1) @@ -42,7 +44,9 @@ def nabla2_on_cell_numpy( def nabla2_on_cell_k_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], psi_c: np.ndarray, geofac_n2s: np.ndarray + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], + psi_c: np.ndarray, + geofac_n2s: np.ndarray, ) -> np.ndarray: c2e2cO = connectivities[dims.C2E2CO] geofac_n2s = np.expand_dims(geofac_n2s, axis=-1) @@ -53,7 +57,7 @@ def nabla2_on_cell_k_numpy( def compute_tangential_wind_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], vn: np.ndarray, rbf_vec_coeff_e: np.ndarray, ) -> np.ndarray: @@ -64,7 +68,7 @@ def compute_tangential_wind_numpy( def interpolate_to_cell_center_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], interpolant: np.ndarray, e_bln_c_s: np.ndarray, **kwargs: Any, @@ -76,7 +80,7 @@ def interpolate_to_cell_center_numpy( def interpolate_cell_field_to_vertex_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], cell_field: np.ndarray, c_intp: np.ndarray, ) -> np.ndarray: @@ -86,7 +90,7 @@ def interpolate_cell_field_to_vertex_numpy( def compute_curl_numpy( - connectivities: Mapping[gtx.FieldOffset, np.ndarray], + connectivities: Mapping[type[gtx.NeighborConnectivity], np.ndarray], edge_field: np.ndarray, geofac_rot: np.ndarray, ) -> np.ndarray: diff --git a/model/testing/src/icon4py/model/testing/stencil_tests.py b/model/testing/src/icon4py/model/testing/stencil_tests.py index dae0d2bb3d..e5f1fe6a3b 100644 --- a/model/testing/src/icon4py/model/testing/stencil_tests.py +++ b/model/testing/src/icon4py/model/testing/stencil_tests.py @@ -260,7 +260,7 @@ class DataAllocationWrapper: grid: base.Grid allocator: gtx_typing.Allocator | None - def connectivity_field(self, offset: str | gtx.FieldOffset) -> gtx.Field: + def connectivity_field(self, offset: type[gtx.NeighborConnectivity]) -> gtx.Field: """ A connectivity table as a regular field, for stencils consuming it as data. @@ -363,13 +363,13 @@ def zero_field( ) -class _NumPyGridConnectivitiesView(Mapping[str | gtx.FieldOffset, np.ndarray]): +class _NumPyGridConnectivitiesView(Mapping[type[gtx.NeighborConnectivity], np.ndarray]): """Read-only `Mapping` exposing a grid's neighbor tables as NumPy arrays.""" def __init__(self, grid: base.Grid) -> None: self._grid = grid - def __getitem__(self, key: str | gtx.FieldOffset) -> np.ndarray: + def __getitem__(self, key: type[gtx.NeighborConnectivity]) -> np.ndarray: # `KeyError` rather than what the grid raises: `Mapping` builds `get` and `in` on # top of this, and both have to see a missing key as missing rather than as an error. try: @@ -380,7 +380,7 @@ def __getitem__(self, key: str | gtx.FieldOffset) -> np.ndarray: raise KeyError(f"Connectivity '{key}' is not a neighbor table.") return connectivity.asnumpy() - def __iter__(self) -> Iterator[str | gtx.FieldOffset]: + def __iter__(self) -> Iterator[type[gtx.NeighborConnectivity]]: return ( key for key, connectivity in self._grid.connectivities.items() @@ -391,16 +391,11 @@ def __len__(self) -> int: return sum(1 for _ in self) -def connectivities_asnumpy(grid: base.Grid) -> Mapping[gtx.FieldOffset, np.ndarray]: - """ - A read-only view of `grid`'s neighbor tables as NumPy arrays. - - Entries can be looked up by `FieldOffset`, as the return annotation advertises, or by - name. The cast is needed because the underlying mapping is keyed by name and `Mapping` - is invariant in its key type, so the honest `Mapping[str | FieldOffset, ...]` would not - be accepted where reference helpers ask for `Mapping[FieldOffset, ...]`. - """ - return cast(Mapping[gtx.FieldOffset, np.ndarray], _NumPyGridConnectivitiesView(grid)) +def connectivities_asnumpy( + grid: base.Grid, +) -> Mapping[type[gtx.NeighborConnectivity], np.ndarray]: + """A read-only view of `grid`'s neighbor tables as NumPy arrays.""" + return _NumPyGridConnectivitiesView(grid) @dataclasses.dataclass(frozen=True) diff --git a/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py b/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py index d26fbe0ad7..4d9616a5cf 100644 --- a/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py +++ b/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py @@ -387,6 +387,10 @@ def test_defaults_select_the_whole_field(self): # -- connectivities_asnumpy ------------------------------------------------------------ +class _Unbound(gtx.NeighborConnectivity[dims.CellDim, dims.EdgeDim]): + class Local(gtx.LocalDimensionIndex): ... + + class StubGrid: """Minimal stand-in exposing only what the connectivities view uses.""" @@ -394,16 +398,15 @@ def __init__(self, connectivities): self.connectivities = connectivities def get_connectivity(self, offset): - return self.connectivities[offset if isinstance(offset, str) else offset.value] + return self.connectivities[offset] class TestConnectivitiesAsNumpy: - def test_lookup_by_name_and_by_field_offset_agree(self, grid): + def test_lookup_matches_the_grid(self, grid): view = stencil_tests.connectivities_asnumpy(grid) assert isinstance(view[dims.E2C], np.ndarray) - np.testing.assert_array_equal(view[dims.E2C], view["E2C"]) - np.testing.assert_array_equal(view[dims.E2C], grid.get_connectivity("E2C").asnumpy()) + np.testing.assert_array_equal(view[dims.E2C], grid.get_connectivity(dims.E2C).asnumpy()) def test_iteration_and_length_cover_the_neighbor_tables(self, grid): view = stencil_tests.connectivities_asnumpy(grid) @@ -417,23 +420,23 @@ def test_iteration_and_length_cover_the_neighbor_tables(self, grid): assert len(view) == len(expected) def test_non_neighbor_table_entries_are_skipped(self, grid): - stub = StubGrid({**dict(grid.connectivities), "Koff": dims.KDim}) + stub = StubGrid({**dict(grid.connectivities), _Unbound: dims.KDim}) view = stencil_tests.connectivities_asnumpy(stub) - assert "Koff" not in set(view) + assert _Unbound not in set(view) assert len(view) == len(set(view)) with pytest.raises(KeyError, match="is not a neighbor table"): - view["Koff"] + view[_Unbound] def test_honours_the_mapping_contract_for_a_missing_key(self, grid): """`get` and `in` are built on `__getitem__`, so it has to raise `KeyError`.""" view = stencil_tests.connectivities_asnumpy(grid) - assert view.get("NoSuchOffset", "default") == "default" - assert "NoSuchOffset" not in view - assert "E2C" in view + assert view.get(_Unbound, "default") == "default" + assert _Unbound not in view + assert dims.E2C in view with pytest.raises(KeyError): - view["NoSuchOffset"] + view[_Unbound] # -- DataAllocationWrapper ------------------------------------------------------------- @@ -505,7 +508,7 @@ def test_connectivity_field_returns_a_plain_field(self, wrapper, grid): assert isinstance(field, gtx.Field) assert not gtx_common.is_neighbor_table(field) - np.testing.assert_array_equal(field.asnumpy(), grid.get_connectivity("E2C").asnumpy()) + np.testing.assert_array_equal(field.asnumpy(), grid.get_connectivity(dims.E2C).asnumpy()) def test_signatures_stay_in_sync_with_data_allocation(self): """The wrapper duplicates the wrapped signatures, so guard against drift.""" From fcd886bf2a204e434b9f4fd412e3393ca1add1c7 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 10:43:03 +0200 Subject: [PATCH 5/6] fix: finish the migration where CI found gaps The first CI run on this branch failed 45 common unit tests and 54 mypy checks. All were icon4py-side: - A dimension's name is `.tag`; `.value` on a dimension class raises. Log messages and the RNG seed in `test_parallel_io` read `.tag`; pytest ids and test messages use `__name__`, since the tag is now a qualified path. `test_icon` reached a connectivity through its local dimension's name and now uses `dim.owner`. mypy does not flag any of these: `value` is declared on the index instances, so reading it on the class type-checks. - `setup_program` annotated `offset_provider` as `gtx_typing.OffsetProvider`, which is the tag-keyed internal form; it now takes `OffsetProviderLike`, what the programs themselves accept. That type is not re-exported from `gt4py.next.typing`, so it comes from `gt4py.next.common`. - There is no string factory for dimensions any more: the test-only dimensions in `test_vertical` are declared classes, and `test_parallel_grid_manager` takes the local dimension from the connectivity it iterates. - `test_halo` indexed the class-keyed neighbor tables by name. - One dimension-keyed dict literal in `test_factory` needed an annotation. `model/common/tests/common` with `--datatest-skip`: 727 passed, 0 failed. mypy: no issues found in 420 source files. --- .../common/decomposition/mpi_decomposition.py | 10 +++++----- .../src/icon4py/model/common/model_options.py | 4 ++-- .../common/decomposition/unit_tests/test_halo.py | 6 +++--- .../grid/mpi_tests/test_parallel_grid_manager.py | 2 +- .../mpi_tests/test_parallel_grid_refinement.py | 4 ++-- .../tests/common/grid/unit_tests/test_icon.py | 7 ++++--- .../common/grid/unit_tests/test_vertical.py | 16 ++++++++++++---- .../common/io/mpi_tests/test_parallel_io.py | 2 +- .../common/states/unit_tests/test_factory.py | 4 +++- 9 files changed, 33 insertions(+), 22 deletions(-) diff --git a/model/common/src/icon4py/model/common/decomposition/mpi_decomposition.py b/model/common/src/icon4py/model/common/decomposition/mpi_decomposition.py index 2988068b17..51b925e28f 100644 --- a/model/common/src/icon4py/model/common/decomposition/mpi_decomposition.py +++ b/model/common/src/icon4py/model/common/decomposition/mpi_decomposition.py @@ -255,7 +255,7 @@ def _create_domain_descriptor(self, dim: gtx.Dimension) -> DomainDescriptor: self._domain_id_gen(), data_alloc.as_numpy(all_global), data_alloc.as_numpy(local_halo) ) log.debug( - f"domain descriptor for dim='{dim.value}' with properties {self._domain_descriptor_info(domain_desc)} created" + f"domain descriptor for dim='{dim.tag}' with properties {self._domain_descriptor_info(domain_desc)} created" ) return domain_desc @@ -266,14 +266,14 @@ def _create_pattern(self, horizontal_dim: gtx.Dimension) -> DomainDescriptor: horizontal_dim, decomp_defs.DecompositionInfo.EntryType.HALO ) halo_generator = HaloGenerator.from_gids(data_alloc.as_numpy(global_halo_idx)) - log.debug(f"halo generator for dim='{horizontal_dim.value}' created") + log.debug(f"halo generator for dim='{horizontal_dim.tag}' created") pattern = make_pattern( self._context, halo_generator, [self._domain_descriptors[horizontal_dim]], ) log.debug( - f"pattern for dim='{horizontal_dim.value}' and {self._domain_descriptor_info(self._domain_descriptors[horizontal_dim])} created" + f"pattern for dim='{horizontal_dim.tag}' and {self._domain_descriptor_info(self._domain_descriptors[horizontal_dim])} created" ) return pattern @@ -336,7 +336,7 @@ def start( patterns=applied_patterns, stream=stream, ) - log.debug(f"exchange for {len(fields)} fields of dimension ='{dim.value}' initiated.") + log.debug(f"exchange for {len(fields)} fields of dimension ='{dim.tag}' initiated.") return MultiNodeResult(handle, applied_patterns) def exchange( @@ -347,7 +347,7 @@ def exchange( ) -> None: # Fall back to the default implementation provided by the protocol. super().exchange(dim, *fields, stream=stream) - log.debug(f"exchange for {len(fields)} fields of dimension ='{dim.value}' done.") + log.debug(f"exchange for {len(fields)} fields of dimension ='{dim.tag}' done.") @dataclass diff --git a/model/common/src/icon4py/model/common/model_options.py b/model/common/src/icon4py/model/common/model_options.py index 11fdc99f6e..ced3b7acba 100644 --- a/model/common/src/icon4py/model/common/model_options.py +++ b/model/common/src/icon4py/model/common/model_options.py @@ -13,7 +13,7 @@ import dace import gt4py.next as gtx import gt4py.next.typing as gtx_typing -from gt4py.next import backend as gtx_backend +from gt4py.next import backend as gtx_backend, common as gtx_common from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations from icon4py.model.common import backend_configuration as backend_cfg, model_backends @@ -154,7 +154,7 @@ def setup_program( variants: dict[str, list[gtx_typing.Scalar]] | None = None, horizontal_sizes: dict[str, gtx.int32] | None = None, vertical_sizes: dict[str, gtx.int32] | None = None, - offset_provider: gtx_typing.OffsetProvider | None = None, + offset_provider: gtx_common.OffsetProviderLike | None = None, backend_config: backend_cfg.BackendConfig | None = None, ) -> Callable[..., None]: """ diff --git a/model/common/tests/common/decomposition/unit_tests/test_halo.py b/model/common/tests/common/decomposition/unit_tests/test_halo.py index 771e46ef14..5973ccecbd 100644 --- a/model/common/tests/common/decomposition/unit_tests/test_halo.py +++ b/model/common/tests/common/decomposition/unit_tests/test_halo.py @@ -84,7 +84,7 @@ def test_halo_constructor_decomposition_info_halo_levels(rank, dim, simple_neigh ) decomp_info = halo_generator(utils.SIMPLE_DISTRIBUTION) my_halo_levels = decomp_info.halo_levels(dim) - print(f"{dim.value}: rank {process_props.rank} has halo levels {my_halo_levels} ") + print(f"{dim.__name__}: rank {process_props.rank} has halo levels {my_halo_levels} ") assert np.all(my_halo_levels != definitions.DecompositionFlag.UNDEFINED), ( "All indices should have a defined DecompositionFlag" ) @@ -161,7 +161,7 @@ def test_no_halo(): def test_halo_constructor_validate_rank_mapping_wrong_shape(simple_neighbor_tables): process_props = utils.DummyProps(rank=2) - num_cells = simple_neighbor_tables["C2E2C"].shape[0] + num_cells = simple_neighbor_tables[dims.C2E2C].shape[0] with pytest.raises(exceptions.ValidationError) as e: halo_generator = halo.IconLikeHaloConstructor( connectivities=simple_neighbor_tables, @@ -174,7 +174,7 @@ def test_halo_constructor_validate_rank_mapping_wrong_shape(simple_neighbor_tabl @pytest.mark.parametrize("rank", (0, 1, 2, 3)) def test_halo_constructor_validate_number_of_node_mismatch(rank, simple_neighbor_tables): process_props = utils.DummyProps(rank=rank) - num_cells = simple_neighbor_tables["C2E2C"].shape[0] + num_cells = simple_neighbor_tables[dims.C2E2C].shape[0] distribution = np.full(num_cells, process_props.comm_size + 1, dtype=int) with pytest.raises(expected_exception=exceptions.ValidationError) as e: halo_generator = halo.IconLikeHaloConstructor( diff --git a/model/common/tests/common/grid/mpi_tests/test_parallel_grid_manager.py b/model/common/tests/common/grid/mpi_tests/test_parallel_grid_manager.py index 64d6951542..481a359559 100644 --- a/model/common/tests/common/grid/mpi_tests/test_parallel_grid_manager.py +++ b/model/common/tests/common/grid/mpi_tests/test_parallel_grid_manager.py @@ -780,7 +780,7 @@ def test_validate_skip_values_in_distributed_connectivities( f"rank={process_props.rank} / {process_props.comm_size}: {k} - # of skip values found in table = {skip_values_in_table}, skip value is {c.skip_value}" ) if skip_values_in_table > 0: - dim = gtx.Dimension(k, gtx.DimensionKind.LOCAL) + dim = gtx_common.local_dimension_of(k) assert ( dim in icon.CONNECTIVITIES_ON_BOUNDARIES or dim in icon.CONNECTIVITIES_ON_PENTAGONS diff --git a/model/common/tests/common/grid/mpi_tests/test_parallel_grid_refinement.py b/model/common/tests/common/grid/mpi_tests/test_parallel_grid_refinement.py index 27f4a59ef5..f81ec6364c 100644 --- a/model/common/tests/common/grid/mpi_tests/test_parallel_grid_refinement.py +++ b/model/common/tests/common/grid/mpi_tests/test_parallel_grid_refinement.py @@ -44,10 +44,10 @@ def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: params = [ (dim, zone) for dim in dims.horizontal_dims() for zone in h_grid._get_zones_for_dim(dim) ] - ids = [f"{dim.value}-{zone}" for dim, zone in params] + ids = [f"{dim.__name__}-{zone}" for dim, zone in params] metafunc.parametrize("dim,zone", params, ids=ids) elif "dim" in metafunc.fixturenames: - ids = [dim.value for dim in dims.horizontal_dims()] + ids = [dim.__name__ for dim in dims.horizontal_dims()] metafunc.parametrize("dim", dims.horizontal_dims(), ids=ids) diff --git a/model/common/tests/common/grid/unit_tests/test_icon.py b/model/common/tests/common/grid/unit_tests/test_icon.py index 20af08a8a3..ae13219ca5 100644 --- a/model/common/tests/common/grid/unit_tests/test_icon.py +++ b/model/common/tests/common/grid/unit_tests/test_icon.py @@ -220,15 +220,16 @@ def test_when_replace_skip_values_then_only_pentagon_points_remain( if dim == dims.LsqUnkDim: pytest.skip("LsqUnkDim is not an offset dimension.") grid = utils.run_grid_manager(grid_description, keep_skip_values=False, backend=backend).grid - connectivity = grid.get_connectivity(dim.value) + assert dim.owner is not None + connectivity = grid.get_connectivity(dim.owner) if dim in icon.CONNECTIVITIES_ON_PENTAGONS and not grid.limited_area: assert np.any(connectivity.asnumpy() == gridfile.GridFile.INVALID_INDEX).item(), ( - f"Connectivity {dim.value} for {grid_description.name} should have skip values." + f"Connectivity {dim.owner.__name__} for {grid_description.name} should have skip values." ) assert connectivity.skip_value == gridfile.GridFile.INVALID_INDEX else: assert not np.any(connectivity.asnumpy() == gridfile.GridFile.INVALID_INDEX).item(), ( - f"Connectivity {dim.value} for {grid_description.name} contains skip values, but none are expected." + f"Connectivity {dim.owner.__name__} for {grid_description.name} contains skip values, but none are expected." ) assert connectivity.skip_value is None diff --git a/model/common/tests/common/grid/unit_tests/test_vertical.py b/model/common/tests/common/grid/unit_tests/test_vertical.py index 2f3c698a20..881b956d83 100644 --- a/model/common/tests/common/grid/unit_tests/test_vertical.py +++ b/model/common/tests/common/grid/unit_tests/test_vertical.py @@ -43,6 +43,15 @@ from icon4py.model.testing import serialbox as sb +class _JDim(gtx.DimensionIndex, kind=gtx.DimensionKind.VERTICAL): ... + + +class _HorizontalDim(gtx.DimensionIndex): ... + + +class _LocalDim(gtx.LocalDimensionIndex): ... + + @pytest.mark.parametrize( "max_h,damping_height,delta,flat_height", [(60000, 34000, 612, 50000), (12000, 10000, 100, 11000), (109050, 45000, 123, 80000)], @@ -117,7 +126,7 @@ def test_grid_size_raises_for_non_vertical_dim( @pytest.mark.datatest def test_grid_size_raises_for_unknown_vertical_dim(grid_savepoint: sb.IconGridSavepoint) -> None: vertical_grid = configure_vertical_grid(grid_savepoint) - j_dim = gtx.Dimension("J", kind=gtx.DimensionKind.VERTICAL) + j_dim = _JDim with pytest.raises(ValueError): vertical_grid.size(j_dim) @@ -178,9 +187,8 @@ def vertical_zones() -> Iterator[v_grid.Zone]: @pytest.mark.parametrize("zone", vertical_zones()) -@pytest.mark.parametrize("kind", (gtx.DimensionKind.LOCAL, gtx.DimensionKind.HORIZONTAL)) -def test_domain_raises_for_non_vertical_dim(zone: v_grid.Zone, kind: gtx.DimensionKind) -> None: - dim = gtx.Dimension("I", kind=kind) +@pytest.mark.parametrize("dim", (_LocalDim, _HorizontalDim), ids=lambda d: d.kind.value) +def test_domain_raises_for_non_vertical_dim(zone: v_grid.Zone, dim: gtx.Dimension) -> None: with pytest.raises(AssertionError): v_grid.Domain(dim, zone) diff --git a/model/common/tests/common/io/mpi_tests/test_parallel_io.py b/model/common/tests/common/io/mpi_tests/test_parallel_io.py index 6f0a958927..5b345510dd 100644 --- a/model/common/tests/common/io/mpi_tests/test_parallel_io.py +++ b/model/common/tests/common/io/mpi_tests/test_parallel_io.py @@ -57,7 +57,7 @@ def synthetic_decomposition_info( info = decomp_defs.DecompositionInfo() for dim, global_size in GLOBAL_SIZES.items(): # deterministic, rank-independent seed (hash() is per-process randomized) - rng = np.random.default_rng(seed=sum(ord(c) for c in dim.value)) + rng = np.random.default_rng(seed=sum(ord(c) for c in dim.tag)) permutation = rng.permutation(global_size) working_ranks = [r for r in range(process_props.comm_size) if r != empty_rank] bounds = np.linspace(0, global_size, len(working_ranks) + 1).astype(int) diff --git a/model/common/tests/common/states/unit_tests/test_factory.py b/model/common/tests/common/states/unit_tests/test_factory.py index a8c09ce8b6..3d8d3951a4 100644 --- a/model/common/tests/common/states/unit_tests/test_factory.py +++ b/model/common/tests/common/states/unit_tests/test_factory.py @@ -163,7 +163,9 @@ def height_coordinate_source( def test_field_operator_provider(cell_coordinate_source: SimpleFieldSource) -> None: field_op = coord_trans.geographical_to_cartesian_on_cells.with_backend(None) - domain = {dims.CellDim: (cell_domain(h_grid.Zone.LOCAL), cell_domain(h_grid.Zone.LOCAL))} + domain: dict[gtx.Dimension, tuple[h_grid.Domain, h_grid.Domain]] = { + dims.CellDim: (cell_domain(h_grid.Zone.LOCAL), cell_domain(h_grid.Zone.LOCAL)) + } deps = {"lat": "lat", "lon": "lon"} fields = {"x": "x", "y": "y", "z": "z"} From 4ac9322572b9307795771c7cfe05a04d45e9d7a3 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 15:36:38 +0200 Subject: [PATCH 6/6] fix: pass the connectivity class to connectivity_field Three test calls still passed `"E2C"` to `connectivity_field`, which now takes the connectivity class; two are tracer-advection stencil tests CI caught, the third is the stencil-test framework's own unit test. The earlier sweep only covered `get_connectivity("...")` and `connectivities["..."]`, and mypy does not check the test trees these live in. No string literal naming a connectivity is left in model/, tools/ or bindings/. --- .../stencil_tests/test_compute_barycentric_backtrajectory.py | 2 +- .../stencil_tests/test_compute_ffsl_backtrajectory.py | 2 +- .../tests/testing/unit_tests/test_stenciltest_framework.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_barycentric_backtrajectory.py b/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_barycentric_backtrajectory.py index 738450df43..3310914bb1 100644 --- a/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_barycentric_backtrajectory.py +++ b/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_barycentric_backtrajectory.py @@ -84,7 +84,7 @@ def reference( def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: p_vn = data_alloc.random_field(dims.EdgeDim, dims.KDim) p_vt = data_alloc.random_field(dims.EdgeDim, dims.KDim) - cell_idx = data_alloc.connectivity_field("E2C") + cell_idx = data_alloc.connectivity_field(dims.E2C) pos_on_tplane_e_1 = data_alloc.random_field(dims.EdgeDim, dims.E2CDim) pos_on_tplane_e_2 = data_alloc.random_field(dims.EdgeDim, dims.E2CDim) primal_normal_cell_1 = data_alloc.random_field(dims.EdgeDim, dims.E2CDim) diff --git a/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_ffsl_backtrajectory.py b/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_ffsl_backtrajectory.py index 18816337c9..3eb3677211 100644 --- a/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_ffsl_backtrajectory.py +++ b/model/atmosphere/tracer_advection/tests/tracer_advection/stencil_tests/test_compute_ffsl_backtrajectory.py @@ -150,7 +150,7 @@ def reference( def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: p_vn = data_alloc.random_field(dims.EdgeDim, dims.KDim) p_vt = data_alloc.random_field(dims.EdgeDim, dims.KDim) - cell_idx = data_alloc.connectivity_field("E2C") + cell_idx = data_alloc.connectivity_field(dims.E2C) cell_blk = data_alloc.constant_field(1, dims.EdgeDim, dims.E2CDim, dtype=gtx.int32) edge_verts_1_x = data_alloc.random_field(dims.EdgeDim) diff --git a/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py b/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py index 4d9616a5cf..4021e79d63 100644 --- a/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py +++ b/model/testing/tests/testing/unit_tests/test_stenciltest_framework.py @@ -504,7 +504,7 @@ def test_connectivity_field_returns_a_plain_field(self, wrapper, grid): A raw `NeighborTable` cannot be passed as a program argument, so stencils that consume a connectivity as data need it re-allocated as an ordinary field. """ - field = wrapper.connectivity_field("E2C") + field = wrapper.connectivity_field(dims.E2C) assert isinstance(field, gtx.Field) assert not gtx_common.is_neighbor_table(field)