From d5249dee691e95447187efbf01a7df244e181908 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 13:20:33 +0200 Subject: [PATCH 01/18] Add the tmx vertical-diffusion operators Tridiagonal matrix assembly for full-level cell, half-level cell and full-level edge fields, the implicit solve (Thomas algorithm as two vertical scans) on the same three grids, and the explicit update on full-level cells. Half-level fields are typed on KHalfDim. Co-authored-by: Jacopo Canton --- .../tmx/stencils/__init__.py | 7 + .../tmx/stencils/vertical_diffusion.py | 252 ++++++++++++++++++ 2 files changed, 259 insertions(+) create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/__init__.py create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/__init__.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/__init__.py new file mode 100644 index 0000000000..de9850de36 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/__init__.py @@ -0,0 +1,7 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py new file mode 100644 index 0000000000..69b5be1a88 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py @@ -0,0 +1,252 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import gt4py.next as gtx +from gt4py.next.experimental import concat_where + +from icon4py.model.common import dimension as dims, field_type_aliases as fa +from icon4py.model.common.type_alias import wpfloat + + +# The forward sweep's init state makes the first row independent of its sub-diagonal entry, +# and the back substitution's init state makes the last row independent of its +# super-diagonal entry. + + +@gtx.scan_operator(axis=dims.KDim, forward=True, init=(wpfloat("0.0"), wpfloat("0.0"))) +def _solve_tridiagonal_matrix_forward_sweep( + state_kminus1: tuple[wpfloat, wpfloat], + a: wpfloat, + b: wpfloat, + c: wpfloat, + d: wpfloat, +) -> tuple[wpfloat, wpfloat]: + c_prime_kminus1, d_prime_kminus1 = state_kminus1 + normalization = wpfloat("1.0") / (b - c_prime_kminus1 * a) + return c * normalization, (d - d_prime_kminus1 * a) * normalization + + +@gtx.scan_operator(axis=dims.KDim, forward=False, init=wpfloat("0.0")) +def _solve_tridiagonal_matrix_back_substitution( + x_kplus1: wpfloat, + c_prime: wpfloat, + d_prime: wpfloat, +) -> wpfloat: + return d_prime - c_prime * x_kplus1 + + +@gtx.scan_operator(axis=dims.KHalfDim, forward=True, init=(wpfloat("0.0"), wpfloat("0.0"))) +def _solve_tridiagonal_matrix_forward_sweep_on_half_levels( + state_kminus1: tuple[wpfloat, wpfloat], + a: wpfloat, + b: wpfloat, + c: wpfloat, + d: wpfloat, +) -> tuple[wpfloat, wpfloat]: + c_prime_kminus1, d_prime_kminus1 = state_kminus1 + normalization = wpfloat("1.0") / (b - c_prime_kminus1 * a) + return c * normalization, (d - d_prime_kminus1 * a) * normalization + + +@gtx.scan_operator(axis=dims.KHalfDim, forward=False, init=wpfloat("0.0")) +def _solve_tridiagonal_matrix_back_substitution_on_half_levels( + x_kplus1: wpfloat, + c_prime: wpfloat, + d_prime: wpfloat, +) -> wpfloat: + return d_prime - c_prime * x_kplus1 + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _assemble_vertical_diffusion_matrix_on_cells( + diffusivity: fa.CellKHalfField[wpfloat], + inv_dz: fa.CellKHalfField[wpfloat], + inv_air_mass: fa.CellKField[wpfloat], + prefactor: wpfloat, + minlvl: gtx.int32, + maxlvl: gtx.int32, +) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: + """ + Sub-, main and super-diagonal of the vertical diffusion matrix for a full-level field. + + The column spans full levels minlvl..maxlvl, with no flux through its top and bottom. + """ + # embedded rejects a scalar branch on an unbounded region, so the zeros are a field + zero = wpfloat("0.0") * inv_air_mass + a = concat_where( + dims.KDim > minlvl, + wpfloat("0.0") + - prefactor * diffusivity(dims.KDim - 0.5) * inv_dz(dims.KDim - 0.5) * inv_air_mass, + zero, + ) + c = concat_where( + dims.KDim < maxlvl, + wpfloat("0.0") + - prefactor * diffusivity(dims.KDim + 0.5) * inv_dz(dims.KDim + 0.5) * inv_air_mass, + zero, + ) + return a, wpfloat("0.0") - a - c, c + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _assemble_vertical_diffusion_matrix_on_cell_half_levels( + diffusivity: fa.CellKField[wpfloat], + inv_dz: fa.CellKField[wpfloat], + inv_air_mass: fa.CellKHalfField[wpfloat], + prefactor: wpfloat, + minlvl: gtx.int32, + maxlvl: gtx.int32, +) -> tuple[fa.CellKHalfField[wpfloat], fa.CellKHalfField[wpfloat], fa.CellKHalfField[wpfloat]]: + """ + Sub-, main and super-diagonal of the vertical diffusion matrix for a half-level field. + + The column spans half levels minlvl..maxlvl, with no flux through its top and bottom. + """ + # embedded rejects a scalar branch on an unbounded region, so the zeros are a field + zero = wpfloat("0.0") * inv_air_mass + a = concat_where( + dims.KHalfDim > minlvl, + wpfloat("0.0") + - prefactor * diffusivity(dims.KHalfDim - 0.5) * inv_dz(dims.KHalfDim - 0.5) * inv_air_mass, + zero, + ) + c = concat_where( + dims.KHalfDim < maxlvl, + wpfloat("0.0") + - prefactor * diffusivity(dims.KHalfDim + 0.5) * inv_dz(dims.KHalfDim + 0.5) * inv_air_mass, + zero, + ) + return a, wpfloat("0.0") - a - c, c + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _assemble_vertical_diffusion_matrix_on_edges( + diffusivity: fa.EdgeKHalfField[wpfloat], + inv_dz: fa.EdgeKHalfField[wpfloat], + inv_air_mass: fa.EdgeKField[wpfloat], + prefactor: wpfloat, + minlvl: gtx.int32, + maxlvl: gtx.int32, +) -> tuple[fa.EdgeKField[wpfloat], fa.EdgeKField[wpfloat], fa.EdgeKField[wpfloat]]: + """ + Sub-, main and super-diagonal of the vertical diffusion matrix for a full-level field. + + The column spans full levels minlvl..maxlvl, with no flux through its top and bottom. + """ + # embedded rejects a scalar branch on an unbounded region, so the zeros are a field + zero = wpfloat("0.0") * inv_air_mass + a = concat_where( + dims.KDim > minlvl, + wpfloat("0.0") + - prefactor * diffusivity(dims.KDim - 0.5) * inv_dz(dims.KDim - 0.5) * inv_air_mass, + zero, + ) + c = concat_where( + dims.KDim < maxlvl, + wpfloat("0.0") + - prefactor * diffusivity(dims.KDim + 0.5) * inv_dz(dims.KDim + 0.5) * inv_air_mass, + zero, + ) + return a, wpfloat("0.0") - a - c, c + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_implicit_vertical_diffusion_on_cells( + var: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + rhs: fa.CellKField[wpfloat], + tend: fa.CellKField[wpfloat], + dtime: wpfloat, +) -> fa.CellKField[wpfloat]: + """ + tend plus the tendency of one implicit step of the vertical diffusion with matrix (a, b, c). + + The system spans the call's vertical domain; a on its first row and c on its last row + have no effect. + """ + inv_dtime = wpfloat("1.0") / dtime + c_prime, d_prime = _solve_tridiagonal_matrix_forward_sweep( + a, inv_dtime + b, c, var * inv_dtime + rhs + ) + return tend + (_solve_tridiagonal_matrix_back_substitution(c_prime, d_prime) - var) * inv_dtime + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_implicit_vertical_diffusion_on_cell_half_levels( + var: fa.CellKHalfField[wpfloat], + a: fa.CellKHalfField[wpfloat], + b: fa.CellKHalfField[wpfloat], + c: fa.CellKHalfField[wpfloat], + rhs: fa.CellKHalfField[wpfloat], + tend: fa.CellKHalfField[wpfloat], + dtime: wpfloat, +) -> fa.CellKHalfField[wpfloat]: + """ + tend plus the tendency of one implicit step of the vertical diffusion with matrix (a, b, c). + + The system spans the call's vertical domain; a on its first row and c on its last row + have no effect. + """ + inv_dtime = wpfloat("1.0") / dtime + c_prime, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels( + a, inv_dtime + b, c, var * inv_dtime + rhs + ) + return ( + tend + + (_solve_tridiagonal_matrix_back_substitution_on_half_levels(c_prime, d_prime) - var) + * inv_dtime + ) + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_implicit_vertical_diffusion_on_edges( + var: fa.EdgeKField[wpfloat], + a: fa.EdgeKField[wpfloat], + b: fa.EdgeKField[wpfloat], + c: fa.EdgeKField[wpfloat], + rhs: fa.EdgeKField[wpfloat], + tend: fa.EdgeKField[wpfloat], + dtime: wpfloat, +) -> fa.EdgeKField[wpfloat]: + """ + tend plus the tendency of one implicit step of the vertical diffusion with matrix (a, b, c). + + The system spans the call's vertical domain; a on its first row and c on its last row + have no effect. + """ + inv_dtime = wpfloat("1.0") / dtime + c_prime, d_prime = _solve_tridiagonal_matrix_forward_sweep( + a, inv_dtime + b, c, var * inv_dtime + rhs + ) + return tend + (_solve_tridiagonal_matrix_back_substitution(c_prime, d_prime) - var) * inv_dtime + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _apply_explicit_vertical_diffusion_on_cells( + var: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + rhs: fa.CellKField[wpfloat], + tend: fa.CellKField[wpfloat], + minlvl: gtx.int32, + maxlvl: gtx.int32, +) -> fa.CellKField[wpfloat]: + """ + tend plus the explicit vertical diffusion tendency rhs - (a, b, c) var. + + The column spans full levels minlvl..maxlvl; a on row minlvl and c on row maxlvl have + no effect. + """ + # embedded rejects a scalar branch on an unbounded region, so the zeros are a field + zero = wpfloat("0.0") * var + from_above = concat_where(dims.KDim > minlvl, a * var(dims.KDim - 1), zero) + from_below = concat_where(dims.KDim < maxlvl, c * var(dims.KDim + 1), zero) + return tend - from_above - b * var - from_below + rhs From a8d135cbabc7beccdb73d24a39f397a3ece18f4a Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 13:20:33 +0200 Subject: [PATCH 02/18] Test the tmx vertical-diffusion operators against numpy --- .../tmx/tests/tmx/stencil_tests/__init__.py | 7 + .../tmx/tests/tmx/stencil_tests/conftest.py | 10 + .../stencil_tests/test_vertical_diffusion.py | 338 ++++++++++++++++++ 3 files changed, 355 insertions(+) create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/__init__.py create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/conftest.py create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/__init__.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/__init__.py new file mode 100644 index 0000000000..de9850de36 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/__init__.py @@ -0,0 +1,7 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/conftest.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/conftest.py new file mode 100644 index 0000000000..03ef56e58d --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/conftest.py @@ -0,0 +1,10 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +from icon4py.model.testing.fixtures.datatest import backend_like +from icon4py.model.testing.fixtures.stencil_tests import data_alloc, grid, grid_manager diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py new file mode 100644 index 0000000000..3e51ab64d1 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py @@ -0,0 +1,338 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause +from typing import Any + +import gt4py.next as gtx +import numpy as np + +from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils.vertical_diffusion import ( + _apply_explicit_vertical_diffusion_on_cells, + _assemble_vertical_diffusion_matrix_on_cell_half_levels, + _assemble_vertical_diffusion_matrix_on_cells, + _assemble_vertical_diffusion_matrix_on_edges, + _solve_implicit_vertical_diffusion_on_cell_half_levels, + _solve_implicit_vertical_diffusion_on_cells, + _solve_implicit_vertical_diffusion_on_edges, +) +from icon4py.model.common import dimension as dims +from icon4py.model.common.grid import base +from icon4py.model.common.type_alias import wpfloat +from icon4py.model.testing import stencil_tests + + +def diffusion_matrix_numpy(interface_coeff: np.ndarray, inv_air_mass: np.ndarray) -> np.ndarray: + """ + Matrix of minus the flux divergence on a column of rows, built column by column from unit + vectors. interface_coeff[:, j] couples rows j and j + 1; no flux crosses the column ends. + """ + num_rows = inv_air_mass.shape[1] + unit_vectors = np.broadcast_to(np.eye(num_rows), (inv_air_mass.shape[0], num_rows, num_rows)) + downward_flux = interface_coeff[:, :, np.newaxis] * (unit_vectors[:, :-1] - unit_vectors[:, 1:]) + downward_flux = np.pad(downward_flux, ((0, 0), (1, 1), (0, 0))) + return inv_air_mass[:, :, np.newaxis] * (downward_flux[:, 1:] - downward_flux[:, :-1]) + + +def tridiagonal_matrix_numpy(a: np.ndarray, b: np.ndarray, c: np.ndarray) -> np.ndarray: + matrix = np.zeros((*b.shape, b.shape[1])) + rows = np.arange(b.shape[1]) + matrix[:, rows, rows] = b + matrix[:, rows[1:], rows[:-1]] = a[:, 1:] + matrix[:, rows[:-1], rows[1:]] = c[:, :-1] + return matrix + + +def matrix_diagonals_on_rows( + matrix: np.ndarray, shape: tuple[int, int], rows: slice +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + a, b, c = np.zeros(shape), np.zeros(shape), np.zeros(shape) + b[:, rows] = np.diagonal(matrix, axis1=1, axis2=2) + a[:, rows][:, 1:] = np.diagonal(matrix, offset=-1, axis1=1, axis2=2) + c[:, rows][:, :-1] = np.diagonal(matrix, offset=1, axis1=1, axis2=2) + return a, b, c + + +def implicit_diffusion_tendency_numpy( + *, + var: np.ndarray, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + rhs: np.ndarray, + tend: np.ndarray, + dtime: float, + rows: slice, +) -> np.ndarray: + matrix = tridiagonal_matrix_numpy(a[:, rows], b[:, rows], c[:, rows]) + matrix += np.eye(matrix.shape[1]) / dtime + new_var = np.linalg.solve(matrix, (var[:, rows] / dtime + rhs[:, rows])[..., np.newaxis]) + out = np.zeros_like(var) + out[:, rows] = tend[:, rows] + (new_var[..., 0] - var[:, rows]) / dtime + return out + + +def vertical_rows(domain: dict[gtx.Dimension, tuple[int, int]], dim: gtx.Dimension) -> slice: + return slice(*domain[dim]) + + +class TestAssembleVerticalDiffusionMatrixOnCells(stencil_tests.StencilTest): + PROGRAM = _assemble_vertical_diffusion_matrix_on_cells + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + diffusivity: np.ndarray, + inv_dz: np.ndarray, + inv_air_mass: np.ndarray, + prefactor: float, + domain: dict, + **kwargs: Any, + ) -> dict: + rows = vertical_rows(domain, dims.KDim) + interfaces = slice(rows.start + 1, rows.stop) + matrix = diffusion_matrix_numpy( + prefactor * diffusivity[:, interfaces] * inv_dz[:, interfaces], + inv_air_mass[:, rows], + ) + return dict(out=matrix_diagonals_on_rows(matrix, inv_air_mass.shape, rows)) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + vertical_start, vertical_end = 1, grid.num_levels - 1 + return dict( + diffusivity=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.0), + inv_dz=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.1), + inv_air_mass=data_alloc.random_field(dims.CellDim, dims.KDim, low=0.1), + prefactor=wpfloat(2.0), + minlvl=gtx.int32(vertical_start), + maxlvl=gtx.int32(vertical_end - 1), + out=tuple(data_alloc.zero_field(dims.CellDim, dims.KDim) for _ in range(3)), + domain={ + dims.CellDim: (0, gtx.int32(grid.num_cells)), + dims.KDim: (gtx.int32(vertical_start), gtx.int32(vertical_end)), + }, + ) + + +class TestAssembleVerticalDiffusionMatrixOnCellHalfLevels(stencil_tests.StencilTest): + PROGRAM = _assemble_vertical_diffusion_matrix_on_cell_half_levels + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + diffusivity: np.ndarray, + inv_dz: np.ndarray, + inv_air_mass: np.ndarray, + prefactor: float, + domain: dict, + **kwargs: Any, + ) -> dict: + rows = vertical_rows(domain, dims.KHalfDim) + interfaces = slice(rows.start, rows.stop - 1) + matrix = diffusion_matrix_numpy( + prefactor * diffusivity[:, interfaces] * inv_dz[:, interfaces], + inv_air_mass[:, rows], + ) + return dict(out=matrix_diagonals_on_rows(matrix, inv_air_mass.shape, rows)) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + vertical_start, vertical_end = 1, grid.num_levels + return dict( + diffusivity=data_alloc.random_field(dims.CellDim, dims.KDim, low=0.0), + inv_dz=data_alloc.random_field(dims.CellDim, dims.KDim, low=0.1), + inv_air_mass=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.1), + prefactor=wpfloat(2.0), + minlvl=gtx.int32(vertical_start), + maxlvl=gtx.int32(vertical_end - 1), + out=tuple(data_alloc.zero_field(dims.CellDim, dims.KHalfDim) for _ in range(3)), + domain={ + dims.CellDim: (0, gtx.int32(grid.num_cells)), + dims.KHalfDim: (gtx.int32(vertical_start), gtx.int32(vertical_end)), + }, + ) + + +class TestAssembleVerticalDiffusionMatrixOnEdges(stencil_tests.StencilTest): + PROGRAM = _assemble_vertical_diffusion_matrix_on_edges + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + diffusivity: np.ndarray, + inv_dz: np.ndarray, + inv_air_mass: np.ndarray, + prefactor: float, + domain: dict, + **kwargs: Any, + ) -> dict: + rows = vertical_rows(domain, dims.KDim) + interfaces = slice(rows.start + 1, rows.stop) + matrix = diffusion_matrix_numpy( + prefactor * diffusivity[:, interfaces] * inv_dz[:, interfaces], + inv_air_mass[:, rows], + ) + return dict(out=matrix_diagonals_on_rows(matrix, inv_air_mass.shape, rows)) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + vertical_start, vertical_end = 0, grid.num_levels + return dict( + diffusivity=data_alloc.random_field(dims.EdgeDim, dims.KHalfDim, low=0.0), + inv_dz=data_alloc.random_field(dims.EdgeDim, dims.KHalfDim, low=0.1), + inv_air_mass=data_alloc.random_field(dims.EdgeDim, dims.KDim, low=0.1), + prefactor=wpfloat(1.0), + minlvl=gtx.int32(vertical_start), + maxlvl=gtx.int32(vertical_end - 1), + out=tuple(data_alloc.zero_field(dims.EdgeDim, dims.KDim) for _ in range(3)), + domain={ + dims.EdgeDim: (0, gtx.int32(grid.num_edges)), + dims.KDim: (gtx.int32(vertical_start), gtx.int32(vertical_end)), + }, + ) + + +def solve_input_data( + data_alloc: stencil_tests.DataAllocationWrapper, + *, + horizontal_dim: gtx.Dimension, + horizontal_size: int, + vertical_dim: gtx.Dimension, + vertical_start: int, + vertical_end: int, +) -> dict: + return dict( + var=data_alloc.random_field(horizontal_dim, vertical_dim), + a=data_alloc.random_field(horizontal_dim, vertical_dim, low=-1.0, high=0.0), + b=data_alloc.random_field(horizontal_dim, vertical_dim, low=2.0, high=3.0), + c=data_alloc.random_field(horizontal_dim, vertical_dim, low=-1.0, high=0.0), + rhs=data_alloc.random_field(horizontal_dim, vertical_dim), + tend=data_alloc.random_field(horizontal_dim, vertical_dim), + dtime=wpfloat(0.5), + out=data_alloc.zero_field(horizontal_dim, vertical_dim), + domain={ + horizontal_dim: (0, gtx.int32(horizontal_size)), + vertical_dim: (gtx.int32(vertical_start), gtx.int32(vertical_end)), + }, + ) + + +class TestSolveImplicitVerticalDiffusionOnCells(stencil_tests.StencilTest): + PROGRAM = _solve_implicit_vertical_diffusion_on_cells + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference(grid: base.Grid, *, domain: dict, out: np.ndarray, **kwargs: Any) -> dict: + return dict( + out=implicit_diffusion_tendency_numpy(**kwargs, rows=vertical_rows(domain, dims.KDim)) + ) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return solve_input_data( + data_alloc, + horizontal_dim=dims.CellDim, + horizontal_size=grid.num_cells, + vertical_dim=dims.KDim, + vertical_start=1, + vertical_end=grid.num_levels - 1, + ) + + +class TestSolveImplicitVerticalDiffusionOnCellHalfLevels(stencil_tests.StencilTest): + PROGRAM = _solve_implicit_vertical_diffusion_on_cell_half_levels + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference(grid: base.Grid, *, domain: dict, out: np.ndarray, **kwargs: Any) -> dict: + return dict( + out=implicit_diffusion_tendency_numpy( + **kwargs, rows=vertical_rows(domain, dims.KHalfDim) + ) + ) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return solve_input_data( + data_alloc, + horizontal_dim=dims.CellDim, + horizontal_size=grid.num_cells, + vertical_dim=dims.KHalfDim, + vertical_start=1, + vertical_end=grid.num_levels, + ) + + +class TestSolveImplicitVerticalDiffusionOnEdges(stencil_tests.StencilTest): + PROGRAM = _solve_implicit_vertical_diffusion_on_edges + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference(grid: base.Grid, *, domain: dict, out: np.ndarray, **kwargs: Any) -> dict: + return dict( + out=implicit_diffusion_tendency_numpy(**kwargs, rows=vertical_rows(domain, dims.KDim)) + ) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return solve_input_data( + data_alloc, + horizontal_dim=dims.EdgeDim, + horizontal_size=grid.num_edges, + vertical_dim=dims.KDim, + vertical_start=0, + vertical_end=grid.num_levels, + ) + + +class TestApplyExplicitVerticalDiffusionOnCells(stencil_tests.StencilTest): + PROGRAM = _apply_explicit_vertical_diffusion_on_cells + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + var: np.ndarray, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + rhs: np.ndarray, + tend: np.ndarray, + domain: dict, + **kwargs: Any, + ) -> dict: + rows = vertical_rows(domain, dims.KDim) + matrix = tridiagonal_matrix_numpy(a[:, rows], b[:, rows], c[:, rows]) + out = np.zeros_like(var) + out[:, rows] = tend[:, rows] + rhs[:, rows] - np.einsum("nij,nj->ni", matrix, var[:, rows]) + return dict(out=out) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + vertical_start, vertical_end = 0, grid.num_levels + return dict( + var=data_alloc.random_field(dims.CellDim, dims.KDim), + a=data_alloc.random_field(dims.CellDim, dims.KDim), + b=data_alloc.random_field(dims.CellDim, dims.KDim), + c=data_alloc.random_field(dims.CellDim, dims.KDim), + rhs=data_alloc.random_field(dims.CellDim, dims.KDim), + tend=data_alloc.random_field(dims.CellDim, dims.KDim), + minlvl=gtx.int32(vertical_start), + maxlvl=gtx.int32(vertical_end - 1), + out=data_alloc.zero_field(dims.CellDim, dims.KDim), + domain={ + dims.CellDim: (0, gtx.int32(grid.num_cells)), + dims.KDim: (gtx.int32(vertical_start), gtx.int32(vertical_end)), + }, + ) From 1470dbc87d534d666e5a26e9c5ac5fab4ffa82bd Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 13:57:33 +0200 Subject: [PATCH 03/18] Move the tridiagonal solver to common and share it with dycore model/common/math/tridiagonal.py holds the Thomas algorithm as a forward sweep and a back substitution scan, in three variants: full levels (wp), half levels (wp) and half levels in mixed precision. The forward sweep carries q = -c' in all of them, as ICON's z_q does. The mixed-precision pair is dycore's former w solver, unchanged; the dycore stencils call it from common. tmx vertical diffusion uses the two wp pairs instead of its own scans. Co-authored-by: Jacopo Canton --- ...diagonal_matrix_for_w_back_substitution.py | 18 +- ..._tridiagonal_matrix_for_w_forward_sweep.py | 38 +--- .../vertically_implicit_dycore_solver.py | 12 +- .../tmx/stencils/vertical_diffusion.py | 71 ++------ .../icon4py/model/common/math/tridiagonal.py | 99 +++++++++++ .../math/stencil_tests/test_tridiagonal.py | 163 ++++++++++++++++++ 6 files changed, 292 insertions(+), 109 deletions(-) create mode 100644 model/common/src/icon4py/model/common/math/tridiagonal.py create mode 100644 model/common/tests/common/math/stencil_tests/test_tridiagonal.py diff --git a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_back_substitution.py b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_back_substitution.py index 1cb3b79398..542987253e 100644 --- a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_back_substitution.py +++ b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_back_substitution.py @@ -6,20 +6,14 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause import gt4py.next as gtx -from gt4py.next import astype from icon4py.model.common import dimension as dims, field_type_aliases as fa +from icon4py.model.common.math.tridiagonal import ( + _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision, +) from icon4py.model.common.type_alias import vpfloat, wpfloat -@gtx.scan_operator(axis=dims.KHalfDim, forward=False, init=wpfloat("0.0")) -def _solve_tridiagonal_matrix_for_w_back_substitution_scan( - w_state: wpfloat, z_q: vpfloat, w: wpfloat -) -> wpfloat: - """Formerly known as _mo_solve_nonhydro_stencil_53_scan.""" - return w + w_state * astype(z_q, wpfloat) # type: ignore[return-value] # return type hints for scan operator broken in GT4Py - - @gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) def solve_tridiagonal_matrix_for_w_back_substitution( z_q: fa.CellKHalfField[vpfloat], @@ -29,9 +23,9 @@ def solve_tridiagonal_matrix_for_w_back_substitution( vertical_start: gtx.int32, vertical_end: gtx.int32, ) -> None: - _solve_tridiagonal_matrix_for_w_back_substitution_scan( - z_q, - w, + _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision( + q=z_q, + d_prime=w, out=w, domain={ dims.CellDim: (horizontal_start, horizontal_end), diff --git a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_forward_sweep.py b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_forward_sweep.py index 7be7eb0977..3659686124 100644 --- a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_forward_sweep.py +++ b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/solve_tridiagonal_matrix_for_w_forward_sweep.py @@ -9,38 +9,10 @@ from gt4py.next import astype from icon4py.model.common import dimension as dims, field_type_aliases as fa -from icon4py.model.common.type_alias import vpfloat, wpfloat - - -@gtx.scan_operator( - axis=dims.KHalfDim, - forward=True, - init=( # type: ignore[call-overload] # GT4Py misses type hint for tuples here - vpfloat("0.0"), - 0.0, - ), # boundary condition for upper tridiagonal element and w at model top +from icon4py.model.common.math.tridiagonal import ( + _solve_tridiagonal_matrix_forward_sweep_on_half_levels_mixed_precision, ) -def tridiagonal_forward_sweep_for_w( - state_kminus1: tuple[vpfloat, float], - a: vpfloat, - b: vpfloat, - c: vpfloat, - d: wpfloat, -) -> tuple[wpfloat, wpfloat]: - """ - | 1 0 | | w_0 | | 0 | | 1 0 | | w_0 | | 0 | - | a_1 b_1 c_1 | | w_1 | | d_1 | | 0 1 cnew_1 | | w_1 | | dnew_1 | - | a_2 b_2 c_2 | | w_2 | = | d_2 | ==> | 0 1 cnew_2 | | w_2 | = | dnew_2 | - | a_3 b_3 c_3 | | w_3 | | d_3 | | 0 1 cnew_3 | | w_3 | | dnew_3 | - | a_4 b_4 c_4 | | w_4 | | d_4 | | 0 1 cnew_4 | | w_4 | | dnew_4 | - | ... | | ... | | ... | | ... | | ... | | ... | - """ - c_kminus1 = astype(state_kminus1[0], vpfloat) - d_kminus1 = state_kminus1[1] - normalization = vpfloat("1.0") / (b + a * c_kminus1) # normalize diagonal element to 1 - c_new = (vpfloat("0.0") - c) * normalization - d_new = (d - astype(a, wpfloat) * d_kminus1) * astype(normalization, wpfloat) - return c_new, d_new # type: ignore[return-value] # return type hints for scan operators broken in GT4Py +from icon4py.model.common.type_alias import vpfloat, wpfloat @gtx.field_operator @@ -68,7 +40,9 @@ def _solve_tridiagonal_matrix_for_w_forward_sweep( w_prep = z_w_expl - z_gamma_wp * ( z_exner_expl(dims.KHalfDim - 0.5) - z_exner_expl(dims.KHalfDim + 0.5) ) - z_q_res, w_res = tridiagonal_forward_sweep_for_w(a=z_a, b=z_b, c=z_c, d=w_prep) + z_q_res, w_res = _solve_tridiagonal_matrix_forward_sweep_on_half_levels_mixed_precision( + a=z_a, b=z_b, c=z_c, d=w_prep + ) return z_q_res, w_res diff --git a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/vertically_implicit_dycore_solver.py b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/vertically_implicit_dycore_solver.py index 94b5a90819..b3c4f16afe 100644 --- a/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/vertically_implicit_dycore_solver.py +++ b/model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/stencils/vertically_implicit_dycore_solver.py @@ -35,9 +35,6 @@ from icon4py.model.atmosphere.dycore.stencils.compute_results_for_thermodynamic_variables import ( _compute_results_for_thermodynamic_variables, ) -from icon4py.model.atmosphere.dycore.stencils.solve_tridiagonal_matrix_for_w_back_substitution import ( - _solve_tridiagonal_matrix_for_w_back_substitution_scan, -) from icon4py.model.atmosphere.dycore.stencils.solve_tridiagonal_matrix_for_w_forward_sweep import ( _solve_tridiagonal_matrix_for_w_forward_sweep, ) @@ -49,6 +46,9 @@ ) from icon4py.model.common import dimension as dims, field_type_aliases as fa, type_alias as ta from icon4py.model.common.constants import PhysicsConstants, RayleighType +from icon4py.model.common.math.tridiagonal import ( + _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision, +) from icon4py.model.common.type_alias import vpfloat, wpfloat @@ -197,9 +197,9 @@ def solve_w( ) next_w = concat_where( dims.KHalfDim < last_inner_level, - _solve_tridiagonal_matrix_for_w_back_substitution_scan( - z_q=tridiagonal_intermediate_result, - w=next_w_intermediate_result, + _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision( + q=tridiagonal_intermediate_result, + d_prime=next_w_intermediate_result, ), next_w, ) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py index 69b5be1a88..6ee6ddb206 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py @@ -10,58 +10,15 @@ from gt4py.next.experimental import concat_where from icon4py.model.common import dimension as dims, field_type_aliases as fa +from icon4py.model.common.math.tridiagonal import ( + _solve_tridiagonal_matrix_back_substitution, + _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp, + _solve_tridiagonal_matrix_forward_sweep, + _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp, +) from icon4py.model.common.type_alias import wpfloat -# The forward sweep's init state makes the first row independent of its sub-diagonal entry, -# and the back substitution's init state makes the last row independent of its -# super-diagonal entry. - - -@gtx.scan_operator(axis=dims.KDim, forward=True, init=(wpfloat("0.0"), wpfloat("0.0"))) -def _solve_tridiagonal_matrix_forward_sweep( - state_kminus1: tuple[wpfloat, wpfloat], - a: wpfloat, - b: wpfloat, - c: wpfloat, - d: wpfloat, -) -> tuple[wpfloat, wpfloat]: - c_prime_kminus1, d_prime_kminus1 = state_kminus1 - normalization = wpfloat("1.0") / (b - c_prime_kminus1 * a) - return c * normalization, (d - d_prime_kminus1 * a) * normalization - - -@gtx.scan_operator(axis=dims.KDim, forward=False, init=wpfloat("0.0")) -def _solve_tridiagonal_matrix_back_substitution( - x_kplus1: wpfloat, - c_prime: wpfloat, - d_prime: wpfloat, -) -> wpfloat: - return d_prime - c_prime * x_kplus1 - - -@gtx.scan_operator(axis=dims.KHalfDim, forward=True, init=(wpfloat("0.0"), wpfloat("0.0"))) -def _solve_tridiagonal_matrix_forward_sweep_on_half_levels( - state_kminus1: tuple[wpfloat, wpfloat], - a: wpfloat, - b: wpfloat, - c: wpfloat, - d: wpfloat, -) -> tuple[wpfloat, wpfloat]: - c_prime_kminus1, d_prime_kminus1 = state_kminus1 - normalization = wpfloat("1.0") / (b - c_prime_kminus1 * a) - return c * normalization, (d - d_prime_kminus1 * a) * normalization - - -@gtx.scan_operator(axis=dims.KHalfDim, forward=False, init=wpfloat("0.0")) -def _solve_tridiagonal_matrix_back_substitution_on_half_levels( - x_kplus1: wpfloat, - c_prime: wpfloat, - d_prime: wpfloat, -) -> wpfloat: - return d_prime - c_prime * x_kplus1 - - @gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) def _assemble_vertical_diffusion_matrix_on_cells( diffusivity: fa.CellKHalfField[wpfloat], @@ -172,10 +129,8 @@ def _solve_implicit_vertical_diffusion_on_cells( have no effect. """ inv_dtime = wpfloat("1.0") / dtime - c_prime, d_prime = _solve_tridiagonal_matrix_forward_sweep( - a, inv_dtime + b, c, var * inv_dtime + rhs - ) - return tend + (_solve_tridiagonal_matrix_back_substitution(c_prime, d_prime) - var) * inv_dtime + q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, inv_dtime + b, c, var * inv_dtime + rhs) + return tend + (_solve_tridiagonal_matrix_back_substitution(q, d_prime) - var) * inv_dtime @gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) @@ -195,12 +150,12 @@ def _solve_implicit_vertical_diffusion_on_cell_half_levels( have no effect. """ inv_dtime = wpfloat("1.0") / dtime - c_prime, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels( + q, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp( a, inv_dtime + b, c, var * inv_dtime + rhs ) return ( tend - + (_solve_tridiagonal_matrix_back_substitution_on_half_levels(c_prime, d_prime) - var) + + (_solve_tridiagonal_matrix_back_substitution_on_half_levels_wp(q, d_prime) - var) * inv_dtime ) @@ -222,10 +177,8 @@ def _solve_implicit_vertical_diffusion_on_edges( have no effect. """ inv_dtime = wpfloat("1.0") / dtime - c_prime, d_prime = _solve_tridiagonal_matrix_forward_sweep( - a, inv_dtime + b, c, var * inv_dtime + rhs - ) - return tend + (_solve_tridiagonal_matrix_back_substitution(c_prime, d_prime) - var) * inv_dtime + q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, inv_dtime + b, c, var * inv_dtime + rhs) + return tend + (_solve_tridiagonal_matrix_back_substitution(q, d_prime) - var) * inv_dtime @gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) diff --git a/model/common/src/icon4py/model/common/math/tridiagonal.py b/model/common/src/icon4py/model/common/math/tridiagonal.py new file mode 100644 index 0000000000..17ab51e094 --- /dev/null +++ b/model/common/src/icon4py/model/common/math/tridiagonal.py @@ -0,0 +1,99 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +""" +Solve the tridiagonal systems ``a_k x_{k-1} + b_k x_k + c_k x_{k+1} = d_k`` along the vertical +with the Thomas algorithm: a forward sweep followed by a back substitution. + +The forward sweep carries ``q = -c'`` (ICON's ``z_q``) and ``d'``. Its init state makes the first row +independent of its sub-diagonal entry, and the back substitution's init state makes the last row +independent of its super-diagonal entry. +""" + +import gt4py.next as gtx +from gt4py.next import astype + +from icon4py.model.common import dimension as dims +from icon4py.model.common.type_alias import vpfloat, wpfloat + + +@gtx.scan_operator(axis=dims.KDim, forward=True, init=(wpfloat("0.0"), wpfloat("0.0"))) +def _solve_tridiagonal_matrix_forward_sweep( + state_kminus1: tuple[wpfloat, wpfloat], + a: wpfloat, + b: wpfloat, + c: wpfloat, + d: wpfloat, +) -> tuple[wpfloat, wpfloat]: + q_kminus1, d_prime_kminus1 = state_kminus1 + normalization = wpfloat("1.0") / (b + a * q_kminus1) + return (wpfloat("0.0") - c) * normalization, (d - a * d_prime_kminus1) * normalization + + +@gtx.scan_operator(axis=dims.KDim, forward=False, init=wpfloat("0.0")) +def _solve_tridiagonal_matrix_back_substitution( + x_kplus1: wpfloat, + q: wpfloat, + d_prime: wpfloat, +) -> wpfloat: + return d_prime + x_kplus1 * q + + +@gtx.scan_operator(axis=dims.KHalfDim, forward=True, init=(wpfloat("0.0"), wpfloat("0.0"))) +def _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp( + state_kminus1: tuple[wpfloat, wpfloat], + a: wpfloat, + b: wpfloat, + c: wpfloat, + d: wpfloat, +) -> tuple[wpfloat, wpfloat]: + q_kminus1, d_prime_kminus1 = state_kminus1 + normalization = wpfloat("1.0") / (b + a * q_kminus1) + return (wpfloat("0.0") - c) * normalization, (d - a * d_prime_kminus1) * normalization + + +@gtx.scan_operator(axis=dims.KHalfDim, forward=False, init=wpfloat("0.0")) +def _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp( + x_kplus1: wpfloat, + q: wpfloat, + d_prime: wpfloat, +) -> wpfloat: + return d_prime + x_kplus1 * q + + +@gtx.scan_operator( + axis=dims.KHalfDim, + forward=True, + init=( # type: ignore[call-overload] # GT4Py misses type hint for tuples here + vpfloat("0.0"), + 0.0, + ), +) +def _solve_tridiagonal_matrix_forward_sweep_on_half_levels_mixed_precision( + state_kminus1: tuple[vpfloat, float], + a: vpfloat, + b: vpfloat, + c: vpfloat, + d: wpfloat, +) -> tuple[wpfloat, wpfloat]: + """The matrix coefficients and ``q`` in vpfloat, ``d'`` in wpfloat.""" + q_kminus1 = astype(state_kminus1[0], vpfloat) + d_prime_kminus1 = state_kminus1[1] + normalization = vpfloat("1.0") / (b + a * q_kminus1) + q = (vpfloat("0.0") - c) * normalization + d_prime = (d - astype(a, wpfloat) * d_prime_kminus1) * astype(normalization, wpfloat) + return q, d_prime # type: ignore[return-value] # return type hints for scan operators broken in GT4Py + + +@gtx.scan_operator(axis=dims.KHalfDim, forward=False, init=wpfloat("0.0")) +def _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision( + x_kplus1: wpfloat, + q: vpfloat, + d_prime: wpfloat, +) -> wpfloat: + return d_prime + x_kplus1 * astype(q, wpfloat) # type: ignore[return-value] # return type hints for scan operator broken in GT4Py diff --git a/model/common/tests/common/math/stencil_tests/test_tridiagonal.py b/model/common/tests/common/math/stencil_tests/test_tridiagonal.py new file mode 100644 index 0000000000..8730ec79d9 --- /dev/null +++ b/model/common/tests/common/math/stencil_tests/test_tridiagonal.py @@ -0,0 +1,163 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause +from typing import Any + +import gt4py.next as gtx +import numpy as np + +from icon4py.model.common import dimension as dims, field_type_aliases as fa +from icon4py.model.common.grid import base +from icon4py.model.common.math.tridiagonal import ( + _solve_tridiagonal_matrix_back_substitution, + _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision, + _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp, + _solve_tridiagonal_matrix_forward_sweep, + _solve_tridiagonal_matrix_forward_sweep_on_half_levels_mixed_precision, + _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp, +) +from icon4py.model.common.type_alias import vpfloat, wpfloat +from icon4py.model.testing import stencil_tests + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_on_full_levels( + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + d: fa.CellKField[wpfloat], +) -> fa.CellKField[wpfloat]: + q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, b, c, d) + return _solve_tridiagonal_matrix_back_substitution(q, d_prime) + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_on_half_levels_wp( + a: fa.CellKHalfField[wpfloat], + b: fa.CellKHalfField[wpfloat], + c: fa.CellKHalfField[wpfloat], + d: fa.CellKHalfField[wpfloat], +) -> fa.CellKHalfField[wpfloat]: + q, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp(a, b, c, d) + return _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp(q, d_prime) + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_on_half_levels_mixed_precision( + a: fa.CellKHalfField[vpfloat], # type: ignore[valid-type] + b: fa.CellKHalfField[vpfloat], # type: ignore[valid-type] + c: fa.CellKHalfField[vpfloat], # type: ignore[valid-type] + d: fa.CellKHalfField[wpfloat], +) -> fa.CellKHalfField[wpfloat]: + q, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels_mixed_precision(a, b, c, d) + return _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision(q, d_prime) + + +def solve_tridiagonal_numpy( + a: np.ndarray, b: np.ndarray, c: np.ndarray, d: np.ndarray +) -> np.ndarray: + """Solve each column's system; a on the first row and c on the last row are not part of it.""" + x = np.empty_like(d) + for cell in range(d.shape[0]): + matrix = np.diag(b[cell]) + np.diag(a[cell, 1:], -1) + np.diag(c[cell, :-1], 1) + x[cell] = np.linalg.solve(matrix, d[cell]) + return x + + +def tridiagonal_input_data( + data_alloc: stencil_tests.DataAllocationWrapper, + grid: base.Grid, + vertical_dim: gtx.Dimension, + coefficient_dtype: type, +) -> dict[str, Any]: + num_rows = grid.num_levels + (1 if vertical_dim == dims.KHalfDim else 0) + # diagonally dominant, so that the system is well conditioned + return dict( + a=data_alloc.random_field( + dims.CellDim, vertical_dim, low=-1.0, high=1.0, dtype=coefficient_dtype + ), + b=data_alloc.random_field( + dims.CellDim, vertical_dim, low=3.0, high=4.0, dtype=coefficient_dtype + ), + c=data_alloc.random_field( + dims.CellDim, vertical_dim, low=-1.0, high=1.0, dtype=coefficient_dtype + ), + d=data_alloc.random_field(dims.CellDim, vertical_dim, dtype=wpfloat), + domain={dims.CellDim: (0, grid.num_cells), vertical_dim: (0, num_rows)}, + out=data_alloc.zero_field(dims.CellDim, vertical_dim, dtype=wpfloat), + ) + + +class TestSolveTridiagonalMatrixOnFullLevels(stencil_tests.StencilTest): + PROGRAM = _solve_on_full_levels + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + d: np.ndarray, + **kwargs: Any, + ) -> dict: + return dict(out=solve_tridiagonal_numpy(a, b, c, d)) + + @stencil_tests.input_data_fixture + def input_data( + data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid + ) -> dict[str, Any]: + return tridiagonal_input_data(data_alloc, grid, dims.KDim, wpfloat) + + +class TestSolveTridiagonalMatrixOnHalfLevelsWp(stencil_tests.StencilTest): + PROGRAM = _solve_on_half_levels_wp + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + d: np.ndarray, + **kwargs: Any, + ) -> dict: + return dict(out=solve_tridiagonal_numpy(a, b, c, d)) + + @stencil_tests.input_data_fixture + def input_data( + data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid + ) -> dict[str, Any]: + return tridiagonal_input_data(data_alloc, grid, dims.KHalfDim, wpfloat) + + +class TestSolveTridiagonalMatrixOnHalfLevelsMixedPrecision(stencil_tests.StencilTest): + PROGRAM = _solve_on_half_levels_mixed_precision + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + d: np.ndarray, + **kwargs: Any, + ) -> dict: + return dict( + out=solve_tridiagonal_numpy(a.astype(wpfloat), b.astype(wpfloat), c.astype(wpfloat), d) + ) + + @stencil_tests.input_data_fixture + def input_data( + data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid + ) -> dict[str, Any]: + return tridiagonal_input_data(data_alloc, grid, dims.KHalfDim, vpfloat) From 126704c4dd1308a75f09c6922087d38e7592e56d Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 14:22:09 +0200 Subject: [PATCH 04/18] Leave the explicit vertical diffusion to the scalar diffusion PR Its only caller is the scalar diffusion's explicit solver. --- .../tmx/stencils/vertical_diffusion.py | 24 ---------- .../stencil_tests/test_vertical_diffusion.py | 44 ------------------- 2 files changed, 68 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py index 6ee6ddb206..9a3c98e74a 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py @@ -179,27 +179,3 @@ def _solve_implicit_vertical_diffusion_on_edges( inv_dtime = wpfloat("1.0") / dtime q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, inv_dtime + b, c, var * inv_dtime + rhs) return tend + (_solve_tridiagonal_matrix_back_substitution(q, d_prime) - var) * inv_dtime - - -@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) -def _apply_explicit_vertical_diffusion_on_cells( - var: fa.CellKField[wpfloat], - a: fa.CellKField[wpfloat], - b: fa.CellKField[wpfloat], - c: fa.CellKField[wpfloat], - rhs: fa.CellKField[wpfloat], - tend: fa.CellKField[wpfloat], - minlvl: gtx.int32, - maxlvl: gtx.int32, -) -> fa.CellKField[wpfloat]: - """ - tend plus the explicit vertical diffusion tendency rhs - (a, b, c) var. - - The column spans full levels minlvl..maxlvl; a on row minlvl and c on row maxlvl have - no effect. - """ - # embedded rejects a scalar branch on an unbounded region, so the zeros are a field - zero = wpfloat("0.0") * var - from_above = concat_where(dims.KDim > minlvl, a * var(dims.KDim - 1), zero) - from_below = concat_where(dims.KDim < maxlvl, c * var(dims.KDim + 1), zero) - return tend - from_above - b * var - from_below + rhs diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py index 3e51ab64d1..f3051f8cb5 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_vertical_diffusion.py @@ -11,7 +11,6 @@ import numpy as np from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils.vertical_diffusion import ( - _apply_explicit_vertical_diffusion_on_cells, _assemble_vertical_diffusion_matrix_on_cell_half_levels, _assemble_vertical_diffusion_matrix_on_cells, _assemble_vertical_diffusion_matrix_on_edges, @@ -293,46 +292,3 @@ def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) vertical_start=0, vertical_end=grid.num_levels, ) - - -class TestApplyExplicitVerticalDiffusionOnCells(stencil_tests.StencilTest): - PROGRAM = _apply_explicit_vertical_diffusion_on_cells - OUTPUTS = ("out",) - - @stencil_tests.static_reference - def reference( - grid: base.Grid, - *, - var: np.ndarray, - a: np.ndarray, - b: np.ndarray, - c: np.ndarray, - rhs: np.ndarray, - tend: np.ndarray, - domain: dict, - **kwargs: Any, - ) -> dict: - rows = vertical_rows(domain, dims.KDim) - matrix = tridiagonal_matrix_numpy(a[:, rows], b[:, rows], c[:, rows]) - out = np.zeros_like(var) - out[:, rows] = tend[:, rows] + rhs[:, rows] - np.einsum("nij,nj->ni", matrix, var[:, rows]) - return dict(out=out) - - @stencil_tests.input_data_fixture - def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: - vertical_start, vertical_end = 0, grid.num_levels - return dict( - var=data_alloc.random_field(dims.CellDim, dims.KDim), - a=data_alloc.random_field(dims.CellDim, dims.KDim), - b=data_alloc.random_field(dims.CellDim, dims.KDim), - c=data_alloc.random_field(dims.CellDim, dims.KDim), - rhs=data_alloc.random_field(dims.CellDim, dims.KDim), - tend=data_alloc.random_field(dims.CellDim, dims.KDim), - minlvl=gtx.int32(vertical_start), - maxlvl=gtx.int32(vertical_end - 1), - out=data_alloc.zero_field(dims.CellDim, dims.KDim), - domain={ - dims.CellDim: (0, gtx.int32(grid.num_cells)), - dims.KDim: (gtx.int32(vertical_start), gtx.int32(vertical_end)), - }, - ) From c178b315ab128158295a4fde3ddf1d1f6418d3fa Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 14:49:30 +0200 Subject: [PATCH 05/18] Add the full tridiagonal solve per field type to common _solve_tridiagonal_matrix_on_{cells,cell_half_levels,edges} run the forward sweep and the back substitution in one call. Field operators are not generic over dimensions, so there is one per field type. The tmx implicit vertical diffusion uses them. --- .../tmx/stencils/vertical_diffusion.py | 25 ++---- .../icon4py/model/common/math/tridiagonal.py | 36 +++++++- .../math/stencil_tests/test_tridiagonal.py | 85 ++++++++++--------- 3 files changed, 88 insertions(+), 58 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py index 9a3c98e74a..a995066813 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/vertical_diffusion.py @@ -11,10 +11,9 @@ from icon4py.model.common import dimension as dims, field_type_aliases as fa from icon4py.model.common.math.tridiagonal import ( - _solve_tridiagonal_matrix_back_substitution, - _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp, - _solve_tridiagonal_matrix_forward_sweep, - _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp, + _solve_tridiagonal_matrix_on_cell_half_levels, + _solve_tridiagonal_matrix_on_cells, + _solve_tridiagonal_matrix_on_edges, ) from icon4py.model.common.type_alias import wpfloat @@ -129,8 +128,8 @@ def _solve_implicit_vertical_diffusion_on_cells( have no effect. """ inv_dtime = wpfloat("1.0") / dtime - q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, inv_dtime + b, c, var * inv_dtime + rhs) - return tend + (_solve_tridiagonal_matrix_back_substitution(q, d_prime) - var) * inv_dtime + x = _solve_tridiagonal_matrix_on_cells(a, inv_dtime + b, c, var * inv_dtime + rhs) + return tend + (x - var) * inv_dtime @gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) @@ -150,14 +149,8 @@ def _solve_implicit_vertical_diffusion_on_cell_half_levels( have no effect. """ inv_dtime = wpfloat("1.0") / dtime - q, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp( - a, inv_dtime + b, c, var * inv_dtime + rhs - ) - return ( - tend - + (_solve_tridiagonal_matrix_back_substitution_on_half_levels_wp(q, d_prime) - var) - * inv_dtime - ) + x = _solve_tridiagonal_matrix_on_cell_half_levels(a, inv_dtime + b, c, var * inv_dtime + rhs) + return tend + (x - var) * inv_dtime @gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) @@ -177,5 +170,5 @@ def _solve_implicit_vertical_diffusion_on_edges( have no effect. """ inv_dtime = wpfloat("1.0") / dtime - q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, inv_dtime + b, c, var * inv_dtime + rhs) - return tend + (_solve_tridiagonal_matrix_back_substitution(q, d_prime) - var) * inv_dtime + x = _solve_tridiagonal_matrix_on_edges(a, inv_dtime + b, c, var * inv_dtime + rhs) + return tend + (x - var) * inv_dtime diff --git a/model/common/src/icon4py/model/common/math/tridiagonal.py b/model/common/src/icon4py/model/common/math/tridiagonal.py index 17ab51e094..9606a13f6d 100644 --- a/model/common/src/icon4py/model/common/math/tridiagonal.py +++ b/model/common/src/icon4py/model/common/math/tridiagonal.py @@ -18,7 +18,7 @@ import gt4py.next as gtx from gt4py.next import astype -from icon4py.model.common import dimension as dims +from icon4py.model.common import dimension as dims, field_type_aliases as fa from icon4py.model.common.type_alias import vpfloat, wpfloat @@ -97,3 +97,37 @@ def _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision( d_prime: wpfloat, ) -> wpfloat: return d_prime + x_kplus1 * astype(q, wpfloat) # type: ignore[return-value] # return type hints for scan operator broken in GT4Py + + +# one per field type: field operators are not generic over dimensions +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_tridiagonal_matrix_on_cells( + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + d: fa.CellKField[wpfloat], +) -> fa.CellKField[wpfloat]: + q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, b, c, d) + return _solve_tridiagonal_matrix_back_substitution(q, d_prime) + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_tridiagonal_matrix_on_cell_half_levels( + a: fa.CellKHalfField[wpfloat], + b: fa.CellKHalfField[wpfloat], + c: fa.CellKHalfField[wpfloat], + d: fa.CellKHalfField[wpfloat], +) -> fa.CellKHalfField[wpfloat]: + q, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp(a, b, c, d) + return _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp(q, d_prime) + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _solve_tridiagonal_matrix_on_edges( + a: fa.EdgeKField[wpfloat], + b: fa.EdgeKField[wpfloat], + c: fa.EdgeKField[wpfloat], + d: fa.EdgeKField[wpfloat], +) -> fa.EdgeKField[wpfloat]: + q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, b, c, d) + return _solve_tridiagonal_matrix_back_substitution(q, d_prime) diff --git a/model/common/tests/common/math/stencil_tests/test_tridiagonal.py b/model/common/tests/common/math/stencil_tests/test_tridiagonal.py index 8730ec79d9..3b18bd7dd7 100644 --- a/model/common/tests/common/math/stencil_tests/test_tridiagonal.py +++ b/model/common/tests/common/math/stencil_tests/test_tridiagonal.py @@ -13,39 +13,16 @@ from icon4py.model.common import dimension as dims, field_type_aliases as fa from icon4py.model.common.grid import base from icon4py.model.common.math.tridiagonal import ( - _solve_tridiagonal_matrix_back_substitution, _solve_tridiagonal_matrix_back_substitution_on_half_levels_mixed_precision, - _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp, - _solve_tridiagonal_matrix_forward_sweep, _solve_tridiagonal_matrix_forward_sweep_on_half_levels_mixed_precision, - _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp, + _solve_tridiagonal_matrix_on_cell_half_levels, + _solve_tridiagonal_matrix_on_cells, + _solve_tridiagonal_matrix_on_edges, ) from icon4py.model.common.type_alias import vpfloat, wpfloat from icon4py.model.testing import stencil_tests -@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) -def _solve_on_full_levels( - a: fa.CellKField[wpfloat], - b: fa.CellKField[wpfloat], - c: fa.CellKField[wpfloat], - d: fa.CellKField[wpfloat], -) -> fa.CellKField[wpfloat]: - q, d_prime = _solve_tridiagonal_matrix_forward_sweep(a, b, c, d) - return _solve_tridiagonal_matrix_back_substitution(q, d_prime) - - -@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) -def _solve_on_half_levels_wp( - a: fa.CellKHalfField[wpfloat], - b: fa.CellKHalfField[wpfloat], - c: fa.CellKHalfField[wpfloat], - d: fa.CellKHalfField[wpfloat], -) -> fa.CellKHalfField[wpfloat]: - q, d_prime = _solve_tridiagonal_matrix_forward_sweep_on_half_levels_wp(a, b, c, d) - return _solve_tridiagonal_matrix_back_substitution_on_half_levels_wp(q, d_prime) - - @gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) def _solve_on_half_levels_mixed_precision( a: fa.CellKHalfField[vpfloat], # type: ignore[valid-type] @@ -71,29 +48,55 @@ def solve_tridiagonal_numpy( def tridiagonal_input_data( data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, + horizontal_dim: gtx.Dimension, vertical_dim: gtx.Dimension, - coefficient_dtype: type, + coefficient_dtype: type = wpfloat, ) -> dict[str, Any]: - num_rows = grid.num_levels + (1 if vertical_dim == dims.KHalfDim else 0) # diagonally dominant, so that the system is well conditioned return dict( a=data_alloc.random_field( - dims.CellDim, vertical_dim, low=-1.0, high=1.0, dtype=coefficient_dtype + horizontal_dim, vertical_dim, low=-1.0, high=1.0, dtype=coefficient_dtype ), b=data_alloc.random_field( - dims.CellDim, vertical_dim, low=3.0, high=4.0, dtype=coefficient_dtype + horizontal_dim, vertical_dim, low=3.0, high=4.0, dtype=coefficient_dtype ), c=data_alloc.random_field( - dims.CellDim, vertical_dim, low=-1.0, high=1.0, dtype=coefficient_dtype + horizontal_dim, vertical_dim, low=-1.0, high=1.0, dtype=coefficient_dtype ), - d=data_alloc.random_field(dims.CellDim, vertical_dim, dtype=wpfloat), - domain={dims.CellDim: (0, grid.num_cells), vertical_dim: (0, num_rows)}, - out=data_alloc.zero_field(dims.CellDim, vertical_dim, dtype=wpfloat), + d=data_alloc.random_field(horizontal_dim, vertical_dim, dtype=wpfloat), + domain={ + horizontal_dim: (0, grid.size[horizontal_dim]), + vertical_dim: (0, grid.size[vertical_dim]), + }, + out=data_alloc.zero_field(horizontal_dim, vertical_dim, dtype=wpfloat), ) -class TestSolveTridiagonalMatrixOnFullLevels(stencil_tests.StencilTest): - PROGRAM = _solve_on_full_levels +class TestSolveTridiagonalMatrixOnCells(stencil_tests.StencilTest): + PROGRAM = _solve_tridiagonal_matrix_on_cells + OUTPUTS = ("out",) + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + d: np.ndarray, + **kwargs: Any, + ) -> dict: + return dict(out=solve_tridiagonal_numpy(a, b, c, d)) + + @stencil_tests.input_data_fixture + def input_data( + data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid + ) -> dict[str, Any]: + return tridiagonal_input_data(data_alloc, grid, dims.CellDim, dims.KDim) + + +class TestSolveTridiagonalMatrixOnEdges(stencil_tests.StencilTest): + PROGRAM = _solve_tridiagonal_matrix_on_edges OUTPUTS = ("out",) @stencil_tests.static_reference @@ -112,11 +115,11 @@ def reference( def input_data( data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid ) -> dict[str, Any]: - return tridiagonal_input_data(data_alloc, grid, dims.KDim, wpfloat) + return tridiagonal_input_data(data_alloc, grid, dims.EdgeDim, dims.KDim) -class TestSolveTridiagonalMatrixOnHalfLevelsWp(stencil_tests.StencilTest): - PROGRAM = _solve_on_half_levels_wp +class TestSolveTridiagonalMatrixOnCellHalfLevels(stencil_tests.StencilTest): + PROGRAM = _solve_tridiagonal_matrix_on_cell_half_levels OUTPUTS = ("out",) @stencil_tests.static_reference @@ -135,7 +138,7 @@ def reference( def input_data( data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid ) -> dict[str, Any]: - return tridiagonal_input_data(data_alloc, grid, dims.KHalfDim, wpfloat) + return tridiagonal_input_data(data_alloc, grid, dims.CellDim, dims.KHalfDim) class TestSolveTridiagonalMatrixOnHalfLevelsMixedPrecision(stencil_tests.StencilTest): @@ -160,4 +163,4 @@ def reference( def input_data( data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid ) -> dict[str, Any]: - return tridiagonal_input_data(data_alloc, grid, dims.KHalfDim, vpfloat) + return tridiagonal_input_data(data_alloc, grid, dims.CellDim, dims.KHalfDim, vpfloat) From 3233b73d987c29f626cbcdaf9b8b24030c7aeea5 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 17:21:52 +0200 Subject: [PATCH 06/18] Add the tmx scalar-diffusion operators Matrix assembly on V's operators, one program per tracer for the implicit vertical and the horizontal diffusion, and the energy path of the temperature diffusion. Co-authored-by: Jacopo Canton --- .../tmx/stencils/scalar_diffusion.py | 408 ++++++++++++++++++ 1 file changed, 408 insertions(+) create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py new file mode 100644 index 0000000000..90b2658db2 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py @@ -0,0 +1,408 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import gt4py.next as gtx +from gt4py.next import broadcast, neighbor_sum +from gt4py.next.experimental import concat_where + +from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils.vertical_diffusion import ( + _assemble_vertical_diffusion_matrix_on_cells, + _solve_implicit_vertical_diffusion_on_cells, +) +from icon4py.model.common import dimension as dims, field_type_aliases as fa +from icon4py.model.common.constants import PhysicsConstants +from icon4py.model.common.dimension import C2E, E2C, C2EDim +from icon4py.model.common.physics.thermodynamics.compute_energy import ( + _compute_dry_static_energy, + compute_internal_energy_per_area, +) +from icon4py.model.common.physics.thermodynamics.compute_temperature import ( + compute_temperature_from_internal_energy_per_area, +) +from icon4py.model.common.type_alias import wpfloat + + +@gtx.field_operator(grid_type=gtx.GridType.UNSTRUCTURED) +def _assemble_scalar_diffusion_matrix( + diffusivity: fa.CellKHalfField[wpfloat], + inv_dz: fa.CellKHalfField[wpfloat], + air_mass: fa.CellKField[wpfloat], + prefactor: wpfloat, + minlvl: gtx.int32, + maxlvl: gtx.int32, +) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: + return _assemble_vertical_diffusion_matrix_on_cells( + diffusivity, inv_dz, wpfloat("1.0") / air_mass, prefactor, minlvl, maxlvl + ) + + +@gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) +def assemble_scalar_diffusion_matrix( + diffusivity: fa.CellKHalfField[wpfloat], + inv_dz: fa.CellKHalfField[wpfloat], + air_mass: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + prefactor: wpfloat, + horizontal_start: gtx.int32, + horizontal_end: gtx.int32, + vertical_start: gtx.int32, + vertical_end: gtx.int32, +) -> None: + _assemble_scalar_diffusion_matrix( + diffusivity, + inv_dz, + air_mass, + prefactor, + vertical_start, + vertical_end - 1, + out=(a, b, c), + domain={ + dims.CellDim: (horizontal_start, horizontal_end), + dims.KDim: (vertical_start, vertical_end), + }, + ) + + +@gtx.field_operator +def _diffuse_scalar( + var: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + rhs: fa.CellKField[wpfloat], + rho: fa.CellKField[wpfloat], + km_ie: fa.EdgeKHalfField[wpfloat], + inv_dual_edge_length: fa.EdgeField[wpfloat], + geofac_div: gtx.Field[gtx.Dims[dims.CellDim, dims.C2EDim], wpfloat], + rturb_prandtl: wpfloat, + prefactor: wpfloat, + dtime: wpfloat, +) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: + """ + New value and tendency of a cell scalar after one diffusion step: implicit in the vertical + with matrix (a, b, c), explicit and conservative in the horizontal. + + var must be valid on the halo cells. + """ + zero = wpfloat("0.0") * var + vertical_tend = _solve_implicit_vertical_diffusion_on_cells(var, a, b, c, rhs, zero, dtime) + flux = ( + wpfloat("0.5") + * prefactor + * rturb_prandtl + * (km_ie(dims.KDim - 0.5) + km_ie(dims.KDim + 0.5)) + * inv_dual_edge_length + * (var(E2C[1]) - var(E2C[0])) + ) + tend = vertical_tend + neighbor_sum(flux(C2E) * geofac_div, axis=C2EDim) / rho + return var + tend * dtime, tend + + +@gtx.field_operator +def _diffuse_tracer( + var: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + surface_flux: fa.CellField[wpfloat], + air_mass: fa.CellKField[wpfloat], + rho: fa.CellKField[wpfloat], + km_ie: fa.EdgeKHalfField[wpfloat], + inv_dual_edge_length: fa.EdgeField[wpfloat], + geofac_div: gtx.Field[gtx.Dims[dims.CellDim, dims.C2EDim], wpfloat], + rturb_prandtl: wpfloat, + prefactor: wpfloat, + dtime: wpfloat, + maxlvl: gtx.int32, +) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: + """:func:`_diffuse_scalar` with the surface flux entering the bottom row maxlvl.""" + rhs = concat_where( + dims.KDim < maxlvl, + wpfloat("0.0") * var, + wpfloat("0.0") - surface_flux * prefactor * (wpfloat("1.0") / air_mass), + ) + return _diffuse_scalar( + var, + a, + b, + c, + rhs, + rho, + km_ie, + inv_dual_edge_length, + geofac_div, + rturb_prandtl, + prefactor, + dtime, + ) + + +@gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) +def diffuse_tracer( + var: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + surface_flux: fa.CellField[wpfloat], + air_mass: fa.CellKField[wpfloat], + rho: fa.CellKField[wpfloat], + km_ie: fa.EdgeKHalfField[wpfloat], + inv_dual_edge_length: fa.EdgeField[wpfloat], + geofac_div: gtx.Field[gtx.Dims[dims.CellDim, dims.C2EDim], wpfloat], + new_var: fa.CellKField[wpfloat], + tend: fa.CellKField[wpfloat], + rturb_prandtl: wpfloat, + prefactor: wpfloat, + dtime: wpfloat, + horizontal_start: gtx.int32, + horizontal_end: gtx.int32, + vertical_start: gtx.int32, + vertical_end: gtx.int32, +) -> None: + _diffuse_tracer( + var, + a, + b, + c, + surface_flux, + air_mass, + rho, + km_ie, + inv_dual_edge_length, + geofac_div, + rturb_prandtl, + prefactor, + dtime, + vertical_end - 1, + out=(new_var, tend), + domain={ + dims.CellDim: (horizontal_start, horizontal_end), + dims.KDim: (vertical_start, vertical_end), + }, + ) + + +@gtx.field_operator +def _compute_energy_from_temperature( + temperature: fa.CellKField[wpfloat], + qv: fa.CellKField[wpfloat], + qc: fa.CellKField[wpfloat], + qi: fa.CellKField[wpfloat], + qr: fa.CellKField[wpfloat], + qs: fa.CellKField[wpfloat], + qg: fa.CellKField[wpfloat], + height_above_ground: fa.CellKField[wpfloat], + grav: wpfloat, + use_internal_energy: bool, +) -> fa.CellKField[wpfloat]: + """ + Specific energy diffused by the heat diffusion: the internal energy plus cvd / cpd times the + geopotential above ground, or the dry static energy. + """ + if use_internal_energy: + one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) + energy = ( + compute_internal_energy_per_area(temperature, qv, qc + qr, qi + qs + qg, one, one) + + grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd + ) + else: + energy = _compute_dry_static_energy(temperature, height_above_ground, grav) + return energy + + +@gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) +def compute_energy_from_temperature( + temperature: fa.CellKField[wpfloat], + qv: fa.CellKField[wpfloat], + qc: fa.CellKField[wpfloat], + qi: fa.CellKField[wpfloat], + qr: fa.CellKField[wpfloat], + qs: fa.CellKField[wpfloat], + qg: fa.CellKField[wpfloat], + height_above_ground: fa.CellKField[wpfloat], + energy: fa.CellKField[wpfloat], + grav: wpfloat, + use_internal_energy: bool, + horizontal_start: gtx.int32, + horizontal_end: gtx.int32, + vertical_start: gtx.int32, + vertical_end: gtx.int32, +) -> None: + _compute_energy_from_temperature( + temperature, + qv, + qc, + qi, + qr, + qs, + qg, + height_above_ground, + grav, + use_internal_energy, + out=energy, + domain={ + dims.CellDim: (horizontal_start, horizontal_end), + dims.KDim: (vertical_start, vertical_end), + }, + ) + + +@gtx.field_operator +def _diffuse_energy_and_update_temperature( + energy: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + sensible_heat_flux: fa.CellField[wpfloat], + evapotranspiration: fa.CellField[wpfloat], + temperature: fa.CellKField[wpfloat], + new_qv: fa.CellKField[wpfloat], + new_qc: fa.CellKField[wpfloat], + new_qi: fa.CellKField[wpfloat], + qr: fa.CellKField[wpfloat], + qs: fa.CellKField[wpfloat], + qg: fa.CellKField[wpfloat], + air_mass: fa.CellKField[wpfloat], + rho: fa.CellKField[wpfloat], + km_ie: fa.EdgeKHalfField[wpfloat], + height_above_ground: fa.CellKField[wpfloat], + inv_dual_edge_length: fa.EdgeField[wpfloat], + geofac_div: gtx.Field[gtx.Dims[dims.CellDim, dims.C2EDim], wpfloat], + rturb_prandtl: wpfloat, + prefactor: wpfloat, + grav: wpfloat, + dtime: wpfloat, + maxlvl: gtx.int32, + use_internal_energy: bool, +) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: + """ + New temperature and its tendency after one diffusion step of the energy of + :func:`_compute_energy_from_temperature`, with the surface energy flux entering the bottom + row maxlvl. + + The new temperature is recovered with the new qv, qc and qi. + """ + inv_air_mass = wpfloat("1.0") / air_mass + # only the bottom row of the surface term is used, where 'temperature' is that of the + # lowest level + if use_internal_energy: + surface_term = ( + wpfloat("0.0") + - ( + sensible_heat_flux + + temperature * evapotranspiration * (PhysicsConstants.cvv - PhysicsConstants.cvd) + ) + * prefactor + * inv_air_mass + ) + else: + surface_term = ( + wpfloat("0.0") + - sensible_heat_flux + * PhysicsConstants.cpd + / PhysicsConstants.cvd + * prefactor + * inv_air_mass + ) + rhs = concat_where(dims.KDim < maxlvl, wpfloat("0.0") * energy, surface_term) + new_energy, _ = _diffuse_scalar( + energy, + a, + b, + c, + rhs, + rho, + km_ie, + inv_dual_edge_length, + geofac_div, + rturb_prandtl, + prefactor, + dtime, + ) + if use_internal_energy: + one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) + new_temperature = compute_temperature_from_internal_energy_per_area( + new_energy - grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd, + new_qv, + new_qc + qr, + new_qi + qs + qg, + one, + one, + ) + else: + new_temperature = (new_energy - grav * height_above_ground) / PhysicsConstants.cpd + return new_temperature, (new_temperature - temperature) * (wpfloat("1.0") / dtime) + + +@gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) +def diffuse_energy_and_update_temperature( + energy: fa.CellKField[wpfloat], + a: fa.CellKField[wpfloat], + b: fa.CellKField[wpfloat], + c: fa.CellKField[wpfloat], + sensible_heat_flux: fa.CellField[wpfloat], + evapotranspiration: fa.CellField[wpfloat], + temperature: fa.CellKField[wpfloat], + new_qv: fa.CellKField[wpfloat], + new_qc: fa.CellKField[wpfloat], + new_qi: fa.CellKField[wpfloat], + qr: fa.CellKField[wpfloat], + qs: fa.CellKField[wpfloat], + qg: fa.CellKField[wpfloat], + air_mass: fa.CellKField[wpfloat], + rho: fa.CellKField[wpfloat], + km_ie: fa.EdgeKHalfField[wpfloat], + height_above_ground: fa.CellKField[wpfloat], + inv_dual_edge_length: fa.EdgeField[wpfloat], + geofac_div: gtx.Field[gtx.Dims[dims.CellDim, dims.C2EDim], wpfloat], + new_temperature: fa.CellKField[wpfloat], + tend_temperature: fa.CellKField[wpfloat], + rturb_prandtl: wpfloat, + prefactor: wpfloat, + grav: wpfloat, + dtime: wpfloat, + use_internal_energy: bool, + horizontal_start: gtx.int32, + horizontal_end: gtx.int32, + vertical_start: gtx.int32, + vertical_end: gtx.int32, +) -> None: + _diffuse_energy_and_update_temperature( + energy, + a, + b, + c, + sensible_heat_flux, + evapotranspiration, + temperature, + new_qv, + new_qc, + new_qi, + qr, + qs, + qg, + air_mass, + rho, + km_ie, + height_above_ground, + inv_dual_edge_length, + geofac_div, + rturb_prandtl, + prefactor, + grav, + dtime, + vertical_end - 1, + use_internal_energy, + out=(new_temperature, tend_temperature), + domain={ + dims.CellDim: (horizontal_start, horizontal_end), + dims.KDim: (vertical_start, vertical_end), + }, + ) From 45f54cefe53f9b771ba668961779e1adac5a1766 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 17:21:52 +0200 Subject: [PATCH 07/18] Add the tmx scalar-diffusion component and its states Co-authored-by: Jacopo Canton --- .../tmx/scalar_diffusion.py | 270 ++++++++++++++++++ .../subgrid_scale_physics/tmx/tmx_states.py | 107 +++++++ 2 files changed, 377 insertions(+) create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py new file mode 100644 index 0000000000..e2da1d09a5 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py @@ -0,0 +1,270 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +"""The scalar diffusion component of tmx. + +Port of ``Compute_diffusion_hydrometeors`` and ``Compute_diffusion_temperature`` in ICON's +``mo_vdf.f90``, with the implicit vertical solver. +""" + +from __future__ import annotations + +import functools +import logging +import typing + +import gt4py.next as gtx + +from icon4py.model.atmosphere.subgrid_scale_physics.tmx import config as tmx_config, tmx_states +from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils import ( + scalar_diffusion as scalar_stencils, +) +from icon4py.model.common import constants, dimension as dims, model_backends +from icon4py.model.common.decomposition import definitions as decomposition +from icon4py.model.common.grid import base as base_grid, horizontal as h_grid +from icon4py.model.common.model_options import setup_program +from icon4py.model.common.utils import data_allocation as data_alloc + + +if typing.TYPE_CHECKING: + import icon4py.model.common.grid.states as grid_states + from icon4py.model.common import field_type_aliases as fa, type_alias as ta + + +log = logging.getLogger(__name__) + + +class ScalarDiffusion: + """The scalar (qv, qc, qi and temperature) diffusion stage of tmx.""" + + def __init__( + self, + *, + grid: base_grid.Grid, + metric_state: tmx_states.TmxMetricState, + interpolation_state: tmx_states.TmxInterpolationState, + edge_params: grid_states.EdgeParams, + backend: model_backends.BackendLike, + exchange: decomposition.ExchangeRuntime, + turb_prandtl: float, + energy_type: tmx_config.EnergyType, + use_scale_turb_energy_flux: bool, + scale_turb_energy_flux: float, + ) -> None: + self._exchange = exchange + # ``zfactor`` in Compute_diffusion_temperature + energy_flux_factor = scale_turb_energy_flux if use_scale_turb_energy_flux else 1.0 + use_internal_energy = energy_type == tmx_config.EnergyType.INTERNAL + + cell_domain = h_grid.domain(dims.CellDim) + cell_start_nudging = grid.start_index(cell_domain(h_grid.Zone.NUDGING)) + cell_end_local = grid.end_index(cell_domain(h_grid.Zone.LOCAL)) + + zero_field = functools.partial( + data_alloc.zero_field, grid, allocator=model_backends.get_allocator(backend) + ) + # the matrix is assembled once for the three tracers and once for the energy + self._matrix_a: fa.CellKField[ta.wpfloat] = zero_field(dims.CellDim, dims.KDim) + self._matrix_b: fa.CellKField[ta.wpfloat] = zero_field(dims.CellDim, dims.KDim) + self._matrix_c: fa.CellKField[ta.wpfloat] = zero_field(dims.CellDim, dims.KDim) + self._zero_surface_flux: fa.CellField[ta.wpfloat] = zero_field(dims.CellDim) + self.energy: fa.CellKField[ta.wpfloat] = zero_field(dims.CellDim, dims.KDim) + + horizontal_sizes = { + "horizontal_start": cell_start_nudging, + "horizontal_end": cell_end_local, + } + vertical_sizes = { + "vertical_start": gtx.int32(0), + "vertical_end": gtx.int32(grid.num_levels), + } + assemble_matrix = functools.partial( + setup_program, + backend=backend, + program=scalar_stencils.assemble_scalar_diffusion_matrix, + horizontal_sizes=horizontal_sizes, + vertical_sizes=vertical_sizes, + offset_provider={}, + ) + self.assemble_tracer_diffusion_matrix = assemble_matrix( + constant_args={"inv_dz": metric_state.inv_ddqz_z_half, "prefactor": 1.0} + ) + self.assemble_energy_diffusion_matrix = assemble_matrix( + constant_args={ + "inv_dz": metric_state.inv_ddqz_z_half, + "prefactor": energy_flux_factor, + } + ) + horizontal_diffusion_args = { + "inv_dual_edge_length": edge_params.inverse_dual_edge_lengths, + "geofac_div": interpolation_state.geofac_div, + "rturb_prandtl": 1.0 / turb_prandtl, + } + self.diffuse_tracer = setup_program( + backend=backend, + program=scalar_stencils.diffuse_tracer, + constant_args={**horizontal_diffusion_args, "prefactor": 1.0}, + horizontal_sizes=horizontal_sizes, + vertical_sizes=vertical_sizes, + offset_provider=grid.connectivities, + ) + self.compute_energy_from_temperature = setup_program( + backend=backend, + program=scalar_stencils.compute_energy_from_temperature, + constant_args={ + "height_above_ground": metric_state.height_above_ground, + "grav": constants.GRAV, + "use_internal_energy": use_internal_energy, + }, + horizontal_sizes=horizontal_sizes, + vertical_sizes=vertical_sizes, + offset_provider={}, + ) + self.diffuse_energy_and_update_temperature = setup_program( + backend=backend, + program=scalar_stencils.diffuse_energy_and_update_temperature, + constant_args={ + **horizontal_diffusion_args, + "height_above_ground": metric_state.height_above_ground, + "prefactor": energy_flux_factor, + "grav": constants.GRAV, + "use_internal_energy": use_internal_energy, + }, + horizontal_sizes=horizontal_sizes, + vertical_sizes=vertical_sizes, + offset_provider=grid.connectivities, + ) + + def run_hydrometeor_diffusion( + self, + *, + input_state: tmx_states.TmxInputState, + surface_flux_state: tmx_states.TmxSurfaceFluxState, + diagnostic_state: tmx_states.TmxDiagnosticState, + tendency_state: tmx_states.TmxTendencyState, + new_state: tmx_states.TmxNewState, + dtime: float, + ) -> None: + """ + Diffuse qv, qc and qi (``Compute_diffusion_hydrometeors`` in mo_vdf.f90, without CO2). + + Only qv has a surface flux, the evapotranspiration. Needs ``kh_ic`` and ``km_ie`` of + ``diagnostic_state``. + """ + log.debug("tmx hydrometeor diffusion (Compute_diffusion_hydrometeors): start") + + log.debug("communication of qv, qc, qi (cells): start") + tracer_exchange = self._exchange.start( + dims.CellDim, input_state.qv, input_state.qc, input_state.qi + ) + + self.assemble_tracer_diffusion_matrix( + diffusivity=diagnostic_state.kh_ic, + air_mass=input_state.air_mass, + a=self._matrix_a, + b=self._matrix_b, + c=self._matrix_c, + ) + + tracer_exchange.finish() + log.debug("communication of qv, qc, qi (cells): end") + + for var, surface_flux, new_var, tend in ( + ( + input_state.qv, + surface_flux_state.evapotranspiration, + new_state.qv, + tendency_state.tend_qv, + ), + (input_state.qc, self._zero_surface_flux, new_state.qc, tendency_state.tend_qc), + (input_state.qi, self._zero_surface_flux, new_state.qi, tendency_state.tend_qi), + ): + self.diffuse_tracer( + var=var, + a=self._matrix_a, + b=self._matrix_b, + c=self._matrix_c, + surface_flux=surface_flux, + air_mass=input_state.air_mass, + rho=input_state.rho, + km_ie=diagnostic_state.km_ie, + new_var=new_var, + tend=tend, + dtime=dtime, + ) + + log.debug("tmx hydrometeor diffusion (Compute_diffusion_hydrometeors): end") + + def run_temperature_diffusion( + self, + *, + input_state: tmx_states.TmxInputState, + surface_flux_state: tmx_states.TmxSurfaceFluxState, + diagnostic_state: tmx_states.TmxDiagnosticState, + tendency_state: tmx_states.TmxTendencyState, + new_state: tmx_states.TmxNewState, + dtime: float, + ) -> None: + """ + Diffuse the temperature as dry static or internal energy (``Compute_diffusion_temperature`` + in mo_vdf.f90). + + The energy is computed with the input tracers and converted back with the new qv, qc and + qi of ``new_state``, so this runs after :meth:`run_hydrometeor_diffusion`. Needs + ``kh_ic`` and ``km_ie`` of ``diagnostic_state``. + """ + log.debug("tmx temperature diffusion (Compute_diffusion_temperature): start") + + self.compute_energy_from_temperature( + temperature=input_state.temperature, + qv=input_state.qv, + qc=input_state.qc, + qi=input_state.qi, + qr=input_state.qr, + qs=input_state.qs, + qg=input_state.qg, + energy=self.energy, + ) + + log.debug("communication of energy (cells): start") + energy_exchange = self._exchange.start(dims.CellDim, self.energy) + + self.assemble_energy_diffusion_matrix( + diffusivity=diagnostic_state.kh_ic, + air_mass=input_state.air_mass, + a=self._matrix_a, + b=self._matrix_b, + c=self._matrix_c, + ) + + energy_exchange.finish() + log.debug("communication of energy (cells): end") + + self.diffuse_energy_and_update_temperature( + energy=self.energy, + a=self._matrix_a, + b=self._matrix_b, + c=self._matrix_c, + sensible_heat_flux=surface_flux_state.sensible_heat_flux, + evapotranspiration=surface_flux_state.evapotranspiration, + temperature=input_state.temperature, + new_qv=new_state.qv, + new_qc=new_state.qc, + new_qi=new_state.qi, + qr=input_state.qr, + qs=input_state.qs, + qg=input_state.qg, + air_mass=input_state.air_mass, + rho=input_state.rho, + km_ie=diagnostic_state.km_ie, + new_temperature=new_state.temperature, + tend_temperature=tendency_state.tend_temperature, + dtime=dtime, + ) + + log.debug("tmx temperature diffusion (Compute_diffusion_temperature): end") diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py index a6d8dc7fbe..713a9bd753 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py @@ -82,6 +82,41 @@ class TmxInterpolationState: """RBF coefficients for the meridional wind component at cell centers (rbf_vec_coeff_c_2).""" +@dataclasses.dataclass(frozen=True) +class TmxSurfaceFluxState: + """Surface fluxes provided by the surface scheme (inputs to the atmospheric diffusion).""" + + evapotranspiration: fa.CellField[ta.wpfloat] + """Surface evapotranspiration flux (``evspsbl``) [kg/(m^2 s)].""" + sensible_heat_flux: fa.CellField[ta.wpfloat] + """Surface sensible heat flux (``hfss``) [W/m^2].""" + u_stress: fa.CellField[ta.wpfloat] + """Zonal surface wind stress (``tauu``) [N/m^2].""" + v_stress: fa.CellField[ta.wpfloat] + """Meridional surface wind stress (``tauv``) [N/m^2].""" + q_snocpymlt: fa.CellField[ta.wpfloat] + """Heating used to melt snow on the canopy [W/m^2].""" + + @classmethod + def allocate( + cls, grid: base_grid.Grid, allocator: gtx_typing.Allocator | None = None + ) -> TmxSurfaceFluxState: + """Allocate a surface flux state with all fields initialized to zero.""" + + def surface(horizontal_dim: gtx.Dimension) -> gtx.Field: + return data_alloc.zero_field( + grid, horizontal_dim, dtype=ta.wpfloat, allocator=allocator + ) + + return cls( + evapotranspiration=surface(dims.CellDim), + sensible_heat_flux=surface(dims.CellDim), + u_stress=surface(dims.CellDim), + v_stress=surface(dims.CellDim), + q_snocpymlt=surface(dims.CellDim), + ) + + @dataclasses.dataclass(frozen=True) class TmxInputState: """Atmospheric input fields of tmx (``t_vdf_atmo_inputs`` in mo_vdf_atmo_memory.f90).""" @@ -98,8 +133,22 @@ class TmxInputState: """Meridional wind (``va``) on full levels [m/s].""" w: fa.CellKHalfField[ta.wpfloat] """Vertical wind (``wa``) on half levels [m/s].""" + qv: fa.CellKField[ta.wpfloat] + """Specific humidity on full levels [kg/kg].""" + qc: fa.CellKField[ta.wpfloat] + """Cloud water mixing ratio on full levels [kg/kg].""" + qi: fa.CellKField[ta.wpfloat] + """Cloud ice mixing ratio on full levels [kg/kg].""" + qr: fa.CellKField[ta.wpfloat] + """Rain mixing ratio on full levels [kg/kg].""" + qs: fa.CellKField[ta.wpfloat] + """Snow mixing ratio on full levels [kg/kg].""" + qg: fa.CellKField[ta.wpfloat] + """Graupel mixing ratio on full levels [kg/kg].""" rho: fa.CellKField[ta.wpfloat] """Air density on full levels [kg/m^3].""" + air_mass: fa.CellKField[ta.wpfloat] + """Air mass per unit area (``mair``) on full levels [kg/m^2].""" @dataclasses.dataclass(frozen=True) @@ -183,3 +232,61 @@ def allocate( w_vert=zero_field(dims.VertexDim, dims.KHalfDim), km_iv=zero_field(dims.VertexDim, dims.KHalfDim), ) + + +@dataclasses.dataclass(frozen=True) +class TmxNewState: + """Fields updated by the tmx diffusion: ``new = state + tend * dtime``.""" + + temperature: fa.CellKField[ta.wpfloat] + """Updated air temperature on full levels [K].""" + qv: fa.CellKField[ta.wpfloat] + """Updated specific humidity on full levels [kg/kg].""" + qc: fa.CellKField[ta.wpfloat] + """Updated cloud water mixing ratio on full levels [kg/kg].""" + qi: fa.CellKField[ta.wpfloat] + """Updated cloud ice mixing ratio on full levels [kg/kg].""" + + @classmethod + def allocate( + cls, grid: base_grid.Grid, allocator: gtx_typing.Allocator | None = None + ) -> TmxNewState: + """Allocate a new state with all fields initialized to zero.""" + zero_field = functools.partial( + data_alloc.zero_field, grid, dtype=ta.wpfloat, allocator=allocator + ) + return cls( + temperature=zero_field(dims.CellDim, dims.KDim), + qv=zero_field(dims.CellDim, dims.KDim), + qc=zero_field(dims.CellDim, dims.KDim), + qi=zero_field(dims.CellDim, dims.KDim), + ) + + +@dataclasses.dataclass(frozen=True) +class TmxTendencyState: + """Tendencies computed by tmx.""" + + tend_temperature: fa.CellKField[ta.wpfloat] + """Air temperature tendency on full levels [K/s].""" + tend_qv: fa.CellKField[ta.wpfloat] + """Specific humidity tendency on full levels [kg/(kg s)].""" + tend_qc: fa.CellKField[ta.wpfloat] + """Cloud water mixing ratio tendency on full levels [kg/(kg s)].""" + tend_qi: fa.CellKField[ta.wpfloat] + """Cloud ice mixing ratio tendency on full levels [kg/(kg s)].""" + + @classmethod + def allocate( + cls, grid: base_grid.Grid, allocator: gtx_typing.Allocator | None = None + ) -> TmxTendencyState: + """Allocate a tendency state with all fields initialized to zero.""" + zero_field = functools.partial( + data_alloc.zero_field, grid, dtype=ta.wpfloat, allocator=allocator + ) + return cls( + tend_temperature=zero_field(dims.CellDim, dims.KDim), + tend_qv=zero_field(dims.CellDim, dims.KDim), + tend_qc=zero_field(dims.CellDim, dims.KDim), + tend_qi=zero_field(dims.CellDim, dims.KDim), + ) From f1ddbda5518df81bfb0ec26428b0e27142e743fa Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 17:21:52 +0200 Subject: [PATCH 08/18] Test the tmx scalar diffusion against numpy and ICON savepoints Co-authored-by: Jacopo Canton --- .../test_tmx_scalar_diffusion.py | 206 +++++++++ .../tmx/tests/tmx/integration_tests/utils.py | 22 + .../stencil_tests/test_scalar_diffusion.py | 437 ++++++++++++++++++ .../src/icon4py/model/testing/serialbox.py | 95 ++++ 4 files changed, 760 insertions(+) create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py create mode 100644 model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py new file mode 100644 index 0000000000..1ed1935755 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py @@ -0,0 +1,206 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +"""Integration tests of the tmx scalar diffusion component. + +Both stages are seeded from the tmx-diagnostics-exit savepoint (``kh_ic``, ``km_ie``) and the +temperature stage also from the tmx-hydro-exit savepoint (the new qv, qc, qi), so that failures +do not cascade between them. +""" + +from __future__ import annotations + +import dataclasses +from typing import TYPE_CHECKING + +import pytest + +from icon4py.model.atmosphere.subgrid_scale_physics.tmx import scalar_diffusion, tmx_states +from icon4py.model.atmosphere.subgrid_scale_physics.tmx.config import TmxConfig +from icon4py.model.common import model_backends +from icon4py.model.common.decomposition import definitions as decomposition +from icon4py.model.testing import definitions, test_utils + +from ..fixtures import * # noqa: F403 +from .utils import ( + RTOL, + TMX_DATES, + TMX_DTIME, + construct_input_state, + construct_interpolation_state, + construct_metric_state, + construct_surface_flux_state, +) + + +if TYPE_CHECKING: + import gt4py.next.typing as gtx_typing + + from icon4py.model.common.grid import icon as icon_grid_ + from icon4py.model.testing import serialbox as sb + + +@dataclasses.dataclass +class _Setup: + component: scalar_diffusion.ScalarDiffusion + input_state: tmx_states.TmxInputState + surface_flux_state: tmx_states.TmxSurfaceFluxState + diagnostic_state: tmx_states.TmxDiagnosticState + tendency_state: tmx_states.TmxTendencyState + new_state: tmx_states.TmxNewState + + +def _setup( + *, + data_provider: sb.IconSerialDataProvider, + grid_savepoint: sb.IconGridSavepoint, + metrics_savepoint: sb.MetricSavepoint, + interpolation_savepoint: sb.InterpolationSavepoint, + icon_grid: icon_grid_.IconGrid, + backend: gtx_typing.Backend | None, + date: str, + config: TmxConfig, +) -> _Setup: + allocator = model_backends.get_allocator(backend) + diagnostics_exit_savepoint = data_provider.from_savepoint_tmx_diagnostics_exit(date=date) + component = scalar_diffusion.ScalarDiffusion( + grid=icon_grid, + metric_state=construct_metric_state( + metrics_savepoint=metrics_savepoint, + init_savepoint=data_provider.from_savepoint_tmx_init(), + allocator=allocator, + ), + interpolation_state=construct_interpolation_state(interpolation_savepoint), + edge_params=grid_savepoint.construct_edge_geometry(), + backend=backend, + exchange=decomposition.SingleNodeExchange(), + turb_prandtl=config.turb_prandtl, + energy_type=config.energy_type, + use_scale_turb_energy_flux=config.use_scale_turb_energy_flux, + scale_turb_energy_flux=config.scale_turb_energy_flux, + ) + return _Setup( + component=component, + input_state=construct_input_state(data_provider.from_savepoint_tmx_entry(date=date)), + surface_flux_state=construct_surface_flux_state( + data_provider.from_savepoint_tmx_surface_fluxes(date=date) + ), + diagnostic_state=dataclasses.replace( + tmx_states.TmxDiagnosticState.allocate(icon_grid, allocator=allocator), + kh_ic=diagnostics_exit_savepoint.kh_ic(), + km_ie=diagnostics_exit_savepoint.km_ie(), + ), + tendency_state=tmx_states.TmxTendencyState.allocate(icon_grid, allocator=allocator), + new_state=tmx_states.TmxNewState.allocate(icon_grid, allocator=allocator), + ) + + +@pytest.mark.datatest +@pytest.mark.parametrize( + "experiment_description, date", + [(definitions.Experiments.EXCLAIM_APE_AES, date) for date in TMX_DATES], +) +def test_tmx_run_hydrometeor_diffusion_single_step( + *, + data_provider: sb.IconSerialDataProvider, + grid_savepoint: sb.IconGridSavepoint, + metrics_savepoint: sb.MetricSavepoint, + interpolation_savepoint: sb.InterpolationSavepoint, + icon_grid: icon_grid_.IconGrid, + backend: gtx_typing.Backend | None, + date: str, + tmx_config: TmxConfig, +) -> None: + setup = _setup( + data_provider=data_provider, + grid_savepoint=grid_savepoint, + metrics_savepoint=metrics_savepoint, + interpolation_savepoint=interpolation_savepoint, + icon_grid=icon_grid, + backend=backend, + date=date, + config=tmx_config, + ) + exit_savepoint = data_provider.from_savepoint_tmx_hydro_exit(date=date) + + setup.component.run_hydrometeor_diffusion( + input_state=setup.input_state, + surface_flux_state=setup.surface_flux_state, + diagnostic_state=setup.diagnostic_state, + tendency_state=setup.tendency_state, + new_state=setup.new_state, + dtime=TMX_DTIME, + ) + + fields = ( + (setup.tendency_state.tend_qv, exit_savepoint.tend_qv(), "tend_qv", 5.0e-20), + (setup.tendency_state.tend_qc, exit_savepoint.tend_qc(), "tend_qc", 5.0e-21), + (setup.tendency_state.tend_qi, exit_savepoint.tend_qi(), "tend_qi", 3.0e-22), + (setup.new_state.qv, exit_savepoint.qv_new(), "qv_new", 2.0e-17), + (setup.new_state.qc, exit_savepoint.qc_new(), "qc_new", 2.0e-18), + (setup.new_state.qi, exit_savepoint.qi_new(), "qi_new", 7.0e-20), + ) + for actual, desired, name, atol in fields: + test_utils.assert_dallclose( + actual.asnumpy(), desired.asnumpy(), rtol=RTOL, atol=atol, err_msg=name + ) + + +@pytest.mark.datatest +@pytest.mark.parametrize( + "experiment_description, date", + [(definitions.Experiments.EXCLAIM_APE_AES, date) for date in TMX_DATES], +) +def test_tmx_run_temperature_diffusion_single_step( + *, + data_provider: sb.IconSerialDataProvider, + grid_savepoint: sb.IconGridSavepoint, + metrics_savepoint: sb.MetricSavepoint, + interpolation_savepoint: sb.InterpolationSavepoint, + icon_grid: icon_grid_.IconGrid, + backend: gtx_typing.Backend | None, + date: str, + tmx_config: TmxConfig, +) -> None: + setup = _setup( + data_provider=data_provider, + grid_savepoint=grid_savepoint, + metrics_savepoint=metrics_savepoint, + interpolation_savepoint=interpolation_savepoint, + icon_grid=icon_grid, + backend=backend, + date=date, + config=tmx_config, + ) + hydro_exit_savepoint = data_provider.from_savepoint_tmx_hydro_exit(date=date) + exit_savepoint = data_provider.from_savepoint_tmx_temperature_exit(date=date) + new_state = dataclasses.replace( + setup.new_state, + qv=hydro_exit_savepoint.qv_new(), + qc=hydro_exit_savepoint.qc_new(), + qi=hydro_exit_savepoint.qi_new(), + ) + + setup.component.run_temperature_diffusion( + input_state=setup.input_state, + surface_flux_state=setup.surface_flux_state, + diagnostic_state=setup.diagnostic_state, + tendency_state=setup.tendency_state, + new_state=new_state, + dtime=TMX_DTIME, + ) + + fields = ( + (setup.component.energy, exit_savepoint.energy(), "energy", 2.0e-10), + (new_state.temperature, exit_savepoint.ta_new(), "ta_new", 4.0e-13), + (setup.tendency_state.tend_temperature, exit_savepoint.tend_ta(), "tend_ta", 2.0e-15), + ) + for actual, desired, name, atol in fields: + test_utils.assert_dallclose( + actual.asnumpy(), desired.asnumpy(), rtol=RTOL, atol=atol, err_msg=name + ) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py index 05901e3fa7..1747246e18 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py @@ -32,6 +32,9 @@ # so the verification tests parametrize over the subsequent steps only. TMX_DATES: tuple[str, ...] = ("2008-09-01T00:05:00.000", "2008-09-01T00:10:00.000") +# 'dt_vdf' of the archive's aes_phy_config [s]. +TMX_DTIME: float = 300.0 + # Relative tolerance of all tmx integration datatests. RTOL: float = 3.0e-12 @@ -96,5 +99,24 @@ def construct_input_state(entry_savepoint: sb.TmxEntrySavepoint) -> tmx_states.T u=entry_savepoint.ua(), v=entry_savepoint.va(), w=entry_savepoint.wa(), + qv=entry_savepoint.qv(), + qc=entry_savepoint.qc(), + qi=entry_savepoint.qi(), + qr=entry_savepoint.qr(), + qs=entry_savepoint.qs(), + qg=entry_savepoint.qg(), rho=entry_savepoint.rho(), + air_mass=entry_savepoint.mair(), + ) + + +def construct_surface_flux_state( + surface_fluxes_savepoint: sb.TmxSurfaceFluxesSavepoint, +) -> tmx_states.TmxSurfaceFluxState: + return tmx_states.TmxSurfaceFluxState( + evapotranspiration=surface_fluxes_savepoint.evspsbl(), + sensible_heat_flux=surface_fluxes_savepoint.hfss(), + u_stress=surface_fluxes_savepoint.tauu(), + v_stress=surface_fluxes_savepoint.tauv(), + q_snocpymlt=surface_fluxes_savepoint.q_snocpymlt(), ) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py new file mode 100644 index 0000000000..b301a6bff1 --- /dev/null +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py @@ -0,0 +1,437 @@ +# ICON4Py - ICON inspired code in Python and GT4Py +# +# Copyright (c) 2022-2024, ETH Zurich and MeteoSwiss +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +from collections.abc import Mapping +from typing import Any + +import gt4py.next as gtx +import numpy as np + +from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils.scalar_diffusion import ( + compute_energy_from_temperature, + diffuse_energy_and_update_temperature, + diffuse_tracer, +) +from icon4py.model.common import constants, dimension as dims +from icon4py.model.common.constants import PhysicsConstants as phy +from icon4py.model.common.grid import base, horizontal as h_grid +from icon4py.model.common.type_alias import wpfloat +from icon4py.model.testing import stencil_tests + +from .test_vertical_diffusion import implicit_diffusion_tendency_numpy + + +_DOMAIN_ARGS = ("horizontal_start", "horizontal_end", "vertical_start", "vertical_end") + + +def moist_heat_capacity_numpy( + qv: np.ndarray, q_liquid: np.ndarray, q_solid: np.ndarray +) -> np.ndarray: + return ( + phy.cvd * (1.0 - qv - q_liquid - q_solid) + + phy.cvv * qv + + phy.cpl * q_liquid + + phy.cpi * q_solid + ) + + +def energy_from_temperature_numpy( + *, + temperature: np.ndarray, + qv: np.ndarray, + q_liquid: np.ndarray, + q_solid: np.ndarray, + height_above_ground: np.ndarray, + grav: float, + use_internal_energy: bool, +) -> np.ndarray: + if use_internal_energy: + return ( + moist_heat_capacity_numpy(qv, q_liquid, q_solid) * temperature + - q_liquid * phy.lvc + - q_solid * phy.lsc + + grav * height_above_ground * phy.cvd / phy.cpd + ) + return phy.cpd * temperature + grav * height_above_ground + + +def temperature_from_energy_numpy( + *, + energy: np.ndarray, + qv: np.ndarray, + q_liquid: np.ndarray, + q_solid: np.ndarray, + height_above_ground: np.ndarray, + grav: float, + use_internal_energy: bool, +) -> np.ndarray: + if use_internal_energy: + internal_energy = energy - grav * height_above_ground * phy.cvd / phy.cpd + return (internal_energy + q_liquid * phy.lvc + q_solid * phy.lsc) / ( + moist_heat_capacity_numpy(qv, q_liquid, q_solid) + ) + return (energy - grav * height_above_ground) / phy.cpd + + +def diffuse_scalar_numpy( + connectivities: Mapping[gtx.FieldOffset, np.ndarray], + *, + var: np.ndarray, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + surface_flux: np.ndarray, + air_mass: np.ndarray, + rho: np.ndarray, + km_ie: np.ndarray, + inv_dual_edge_length: np.ndarray, + geofac_div: np.ndarray, + rturb_prandtl: float, + prefactor: float, + dtime: float, + cells: slice, + rows: slice, +) -> tuple[np.ndarray, np.ndarray]: + """New value and tendency on cells x rows; surface_flux is per cell of the bottom row.""" + rhs = np.zeros_like(var) + bottom = rows.stop - 1 + rhs[:, bottom] = -surface_flux * prefactor / air_mass[:, bottom] + vertical_tend = implicit_diffusion_tendency_numpy( + var=var, a=a, b=b, c=c, rhs=rhs, tend=np.zeros_like(var), dtime=dtime, rows=rows + ) + + e2c = connectivities[dims.E2C] + c2e = connectivities[dims.C2E] + diffusivity_e = prefactor * rturb_prandtl * 0.5 * (km_ie[:, :-1] + km_ie[:, 1:]) + flux_e = diffusivity_e * inv_dual_edge_length[:, np.newaxis] * (var[e2c[:, 1]] - var[e2c[:, 0]]) + divergence = np.sum(flux_e[c2e] * geofac_div[:, :, np.newaxis], axis=1) + + tend = np.zeros_like(var) + new_var = np.zeros_like(var) + tend[cells, rows] = vertical_tend[cells, rows] + divergence[cells, rows] / rho[cells, rows] + new_var[cells, rows] = var[cells, rows] + tend[cells, rows] * dtime + return new_var, tend + + +def _cells(grid: base.Grid) -> tuple[gtx.int32, gtx.int32]: + cell_domain = h_grid.domain(dims.CellDim) + return ( + grid.start_index(cell_domain(h_grid.Zone.NUDGING)), + grid.end_index(cell_domain(h_grid.Zone.LOCAL)), + ) + + +def _diffusion_input_data( + data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, var_name: str +) -> dict: + horizontal_start, horizontal_end = _cells(grid) + return { + var_name: data_alloc.random_field(dims.CellDim, dims.KDim), + "a": data_alloc.random_field(dims.CellDim, dims.KDim, low=-1.0, high=0.0), + "b": data_alloc.random_field(dims.CellDim, dims.KDim, low=2.0, high=3.0), + "c": data_alloc.random_field(dims.CellDim, dims.KDim, low=-1.0, high=0.0), + "air_mass": data_alloc.random_field(dims.CellDim, dims.KDim, low=1.0, high=2.0), + "rho": data_alloc.random_field(dims.CellDim, dims.KDim, low=0.5, high=1.5), + "km_ie": data_alloc.random_field(dims.EdgeDim, dims.KHalfDim, low=0.0), + "inv_dual_edge_length": data_alloc.random_field(dims.EdgeDim, low=0.1), + "geofac_div": data_alloc.random_field(dims.CellDim, dims.C2EDim), + "rturb_prandtl": wpfloat(3.0), + "dtime": wpfloat(0.5), + "horizontal_start": horizontal_start, + "horizontal_end": horizontal_end, + "vertical_start": gtx.int32(0), + "vertical_end": gtx.int32(grid.num_levels), + } + + +def _slices(horizontal_start, horizontal_end, vertical_start, vertical_end) -> tuple[slice, slice]: + return slice(horizontal_start, horizontal_end), slice(vertical_start, vertical_end) + + +class TestDiffuseTracer(stencil_tests.StencilTest): + PROGRAM = diffuse_tracer + OUTPUTS = ("new_var", "tend") + STATIC_PARAMS = { + stencil_tests.StandardStaticVariants.NONE: (), + stencil_tests.StandardStaticVariants.COMPILE_TIME_DOMAIN: ( + *_DOMAIN_ARGS, + "rturb_prandtl", + "prefactor", + ), + } + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + var: np.ndarray, + a: np.ndarray, + b: np.ndarray, + c: np.ndarray, + surface_flux: np.ndarray, + air_mass: np.ndarray, + rho: np.ndarray, + km_ie: np.ndarray, + inv_dual_edge_length: np.ndarray, + geofac_div: np.ndarray, + rturb_prandtl: float, + prefactor: float, + dtime: float, + horizontal_start: int, + horizontal_end: int, + vertical_start: int, + vertical_end: int, + **kwargs: Any, + ) -> dict: + cells, rows = _slices(horizontal_start, horizontal_end, vertical_start, vertical_end) + new_var, tend = diffuse_scalar_numpy( + stencil_tests.connectivities_asnumpy(grid), + var=var, + a=a, + b=b, + c=c, + surface_flux=surface_flux, + air_mass=air_mass, + rho=rho, + km_ie=km_ie, + inv_dual_edge_length=inv_dual_edge_length, + geofac_div=geofac_div, + rturb_prandtl=rturb_prandtl, + prefactor=prefactor, + dtime=dtime, + cells=cells, + rows=rows, + ) + return dict(new_var=new_var, tend=tend) + + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return dict( + **_diffusion_input_data(data_alloc, grid, "var"), + surface_flux=data_alloc.random_field(dims.CellDim), + new_var=data_alloc.zero_field(dims.CellDim, dims.KDim), + tend=data_alloc.zero_field(dims.CellDim, dims.KDim), + prefactor=wpfloat(1.5), + ) + + +def _q_liquid_and_solid(qc, qi, qr, qs, qg) -> tuple[np.ndarray, np.ndarray]: + return qc + qr, qi + qs + qg + + +def _tracers(data_alloc: stencil_tests.DataAllocationWrapper, prefix: str = "") -> dict: + return { + f"{prefix}{name}": data_alloc.random_field(dims.CellDim, dims.KDim, low=0.0, high=0.01) + for name in ("qv", "qc", "qi") + } | { + name: data_alloc.random_field(dims.CellDim, dims.KDim, low=0.0, high=0.01) + for name in ("qr", "qs", "qg") + } + + +class _ComputeEnergyFromTemperature: + PROGRAM = compute_energy_from_temperature + OUTPUTS = ("energy",) + STATIC_PARAMS = { + stencil_tests.StandardStaticVariants.NONE: (), + stencil_tests.StandardStaticVariants.COMPILE_TIME_DOMAIN: ( + *_DOMAIN_ARGS, + "grav", + "use_internal_energy", + ), + } + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + temperature: np.ndarray, + qv: np.ndarray, + qc: np.ndarray, + qi: np.ndarray, + qr: np.ndarray, + qs: np.ndarray, + qg: np.ndarray, + height_above_ground: np.ndarray, + grav: float, + use_internal_energy: bool, + horizontal_start: int, + horizontal_end: int, + vertical_start: int, + vertical_end: int, + **kwargs: Any, + ) -> dict: + cells, rows = _slices(horizontal_start, horizontal_end, vertical_start, vertical_end) + q_liquid, q_solid = _q_liquid_and_solid(qc, qi, qr, qs, qg) + energy = np.zeros_like(temperature) + energy[cells, rows] = energy_from_temperature_numpy( + temperature=temperature, + qv=qv, + q_liquid=q_liquid, + q_solid=q_solid, + height_above_ground=height_above_ground, + grav=grav, + use_internal_energy=use_internal_energy, + )[cells, rows] + return dict(energy=energy) + + +def _compute_energy_input_data( + data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, use_internal_energy: bool +) -> dict: + horizontal_start, horizontal_end = _cells(grid) + return dict( + temperature=data_alloc.random_field(dims.CellDim, dims.KDim, low=200.0, high=300.0), + **_tracers(data_alloc), + height_above_ground=data_alloc.random_field(dims.CellDim, dims.KDim, high=1.0e4), + energy=data_alloc.zero_field(dims.CellDim, dims.KDim), + grav=constants.GRAV, + use_internal_energy=use_internal_energy, + horizontal_start=horizontal_start, + horizontal_end=horizontal_end, + vertical_start=gtx.int32(0), + vertical_end=gtx.int32(grid.num_levels), + ) + + +class TestComputeInternalEnergyFromTemperature( + _ComputeEnergyFromTemperature, stencil_tests.StencilTest +): + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return _compute_energy_input_data(data_alloc, grid, use_internal_energy=True) + + +class TestComputeDryStaticEnergyFromTemperature( + _ComputeEnergyFromTemperature, stencil_tests.StencilTest +): + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return _compute_energy_input_data(data_alloc, grid, use_internal_energy=False) + + +class _DiffuseEnergyAndUpdateTemperature: + PROGRAM = diffuse_energy_and_update_temperature + OUTPUTS = ("new_temperature", "tend_temperature") + STATIC_PARAMS = { + stencil_tests.StandardStaticVariants.NONE: (), + stencil_tests.StandardStaticVariants.COMPILE_TIME_DOMAIN: ( + *_DOMAIN_ARGS, + "rturb_prandtl", + "prefactor", + "grav", + "use_internal_energy", + ), + } + + @stencil_tests.static_reference + def reference( + grid: base.Grid, + *, + energy: np.ndarray, + sensible_heat_flux: np.ndarray, + evapotranspiration: np.ndarray, + temperature: np.ndarray, + new_qv: np.ndarray, + new_qc: np.ndarray, + new_qi: np.ndarray, + qr: np.ndarray, + qs: np.ndarray, + qg: np.ndarray, + height_above_ground: np.ndarray, + prefactor: float, + grav: float, + dtime: float, + use_internal_energy: bool, + horizontal_start: int, + horizontal_end: int, + vertical_start: int, + vertical_end: int, + **kwargs: Any, + ) -> dict: + cells, rows = _slices(horizontal_start, horizontal_end, vertical_start, vertical_end) + if use_internal_energy: + temperature_sfc = temperature[:, vertical_end - 1] + surface_flux = sensible_heat_flux + temperature_sfc * evapotranspiration * ( + phy.cvv - phy.cvd + ) + else: + surface_flux = sensible_heat_flux * phy.cpd / phy.cvd + new_energy, _ = diffuse_scalar_numpy( + stencil_tests.connectivities_asnumpy(grid), + var=energy, + surface_flux=surface_flux, + prefactor=prefactor, + dtime=dtime, + cells=cells, + rows=rows, + **{ + name: kwargs[name] + for name in ( + "a", + "b", + "c", + "air_mass", + "rho", + "km_ie", + "inv_dual_edge_length", + "geofac_div", + "rturb_prandtl", + ) + }, + ) + q_liquid, q_solid = _q_liquid_and_solid(new_qc, new_qi, qr, qs, qg) + new_temperature = np.zeros_like(temperature) + tend_temperature = np.zeros_like(temperature) + new_temperature[cells, rows] = temperature_from_energy_numpy( + energy=new_energy, + qv=new_qv, + q_liquid=q_liquid, + q_solid=q_solid, + height_above_ground=height_above_ground, + grav=grav, + use_internal_energy=use_internal_energy, + )[cells, rows] + tend_temperature[cells, rows] = ( + new_temperature[cells, rows] - temperature[cells, rows] + ) / dtime + return dict(new_temperature=new_temperature, tend_temperature=tend_temperature) + + +def _diffuse_energy_input_data( + data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, use_internal_energy: bool +) -> dict: + return dict( + **_diffusion_input_data(data_alloc, grid, "energy"), + **_tracers(data_alloc, prefix="new_"), + sensible_heat_flux=data_alloc.random_field(dims.CellDim), + evapotranspiration=data_alloc.random_field(dims.CellDim), + temperature=data_alloc.random_field(dims.CellDim, dims.KDim, low=200.0, high=300.0), + height_above_ground=data_alloc.random_field(dims.CellDim, dims.KDim, high=1.0e4), + new_temperature=data_alloc.zero_field(dims.CellDim, dims.KDim), + tend_temperature=data_alloc.zero_field(dims.CellDim, dims.KDim), + prefactor=wpfloat(1.5), + grav=constants.GRAV, + use_internal_energy=use_internal_energy, + ) + + +class TestDiffuseInternalEnergyAndUpdateTemperature( + _DiffuseEnergyAndUpdateTemperature, stencil_tests.StencilTest +): + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return _diffuse_energy_input_data(data_alloc, grid, use_internal_energy=True) + + +class TestDiffuseDryStaticEnergyAndUpdateTemperature( + _DiffuseEnergyAndUpdateTemperature, stencil_tests.StencilTest +): + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return _diffuse_energy_input_data(data_alloc, grid, use_internal_energy=False) diff --git a/model/testing/src/icon4py/model/testing/serialbox.py b/model/testing/src/icon4py/model/testing/serialbox.py index 8c386b7f25..398165b3e8 100644 --- a/model/testing/src/icon4py/model/testing/serialbox.py +++ b/model/testing/src/icon4py/model/testing/serialbox.py @@ -2020,6 +2020,24 @@ def va(self): def wa(self): return self._get_field("wa", dims.CellDim, dims.KHalfDim) + def qv(self): + return self._get_field("qv", dims.CellDim, dims.KDim) + + def qc(self): + return self._get_field("qc", dims.CellDim, dims.KDim) + + def qi(self): + return self._get_field("qi", dims.CellDim, dims.KDim) + + def qr(self): + return self._get_field("qr", dims.CellDim, dims.KDim) + + def qs(self): + return self._get_field("qs", dims.CellDim, dims.KDim) + + def qg(self): + return self._get_field("qg", dims.CellDim, dims.KDim) + def rho(self): return self._get_field("rho", dims.CellDim, dims.KDim) @@ -2029,6 +2047,28 @@ def tempv(self): def pres(self): return self._get_field("pres", dims.CellDim, dims.KDim) + def mair(self): + return self._get_field("mair", dims.CellDim, dims.KDim) + + +class TmxSurfaceFluxesSavepoint(IconSavepoint): + """Savepoint after the surface model call in vdf Compute in mo_vdf.f90.""" + + def evspsbl(self): + return self._get_field("evspsbl", dims.CellDim) + + def hfss(self): + return self._get_field("hfss", dims.CellDim) + + def tauu(self): + return self._get_field("tauu", dims.CellDim) + + def tauv(self): + return self._get_field("tauv", dims.CellDim) + + def q_snocpymlt(self): + return self._get_field("q_snocpymlt", dims.CellDim) + class TmxDiagnosticsExitSavepoint(IconSavepoint): """Savepoint at exit of vdf Compute_diagnostics in mo_vdf_atmo.f90.""" @@ -2097,6 +2137,41 @@ def km_ie(self): return self._get_field("km_ie", dims.EdgeDim, dims.KHalfDim) +class TmxHydroExitSavepoint(IconSavepoint): + """Savepoint after Compute_diffusion_hydrometeors in mo_vdf.f90.""" + + def tend_qv(self): + return self._get_field("tend_qv", dims.CellDim, dims.KDim) + + def tend_qc(self): + return self._get_field("tend_qc", dims.CellDim, dims.KDim) + + def tend_qi(self): + return self._get_field("tend_qi", dims.CellDim, dims.KDim) + + def qv_new(self): + return self._get_field("qv_new", dims.CellDim, dims.KDim) + + def qc_new(self): + return self._get_field("qc_new", dims.CellDim, dims.KDim) + + def qi_new(self): + return self._get_field("qi_new", dims.CellDim, dims.KDim) + + +class TmxTemperatureExitSavepoint(IconSavepoint): + """Savepoint at exit of Compute_diffusion_temperature in mo_vdf.f90.""" + + def energy(self): + return self._get_field("energy", dims.CellDim, dims.KDim) + + def tend_ta(self): + return self._get_field("tend_ta", dims.CellDim, dims.KDim) + + def ta_new(self): + return self._get_field("ta_new", dims.CellDim, dims.KDim) + + class IconTimeStepExitSavepoint(IconSavepoint): """End-of-timestep prognostic state, written in perform_nh_timeloop right after integrate_nh returns: all physics tendencies applied, time levels swapped.""" @@ -2531,6 +2606,12 @@ def from_savepoint_tmx_entry(self, date: str) -> TmxEntrySavepoint: savepoint, self.serializer, size=self.grid_size, backend=self.backend ) + def from_savepoint_tmx_surface_fluxes(self, date: str) -> TmxSurfaceFluxesSavepoint: + savepoint = self.serializer.savepoint["tmx-surface-fluxes"].id[1].date[date].as_savepoint() + return TmxSurfaceFluxesSavepoint( + savepoint, self.serializer, size=self.grid_size, backend=self.backend + ) + def from_savepoint_tmx_diagnostics_exit(self, date: str) -> TmxDiagnosticsExitSavepoint: savepoint = ( self.serializer.savepoint["tmx-diagnostics-exit"].id[1].date[date].as_savepoint() @@ -2538,3 +2619,17 @@ def from_savepoint_tmx_diagnostics_exit(self, date: str) -> TmxDiagnosticsExitSa return TmxDiagnosticsExitSavepoint( savepoint, self.serializer, size=self.grid_size, backend=self.backend ) + + def from_savepoint_tmx_hydro_exit(self, date: str) -> TmxHydroExitSavepoint: + savepoint = self.serializer.savepoint["tmx-hydro-exit"].id[1].date[date].as_savepoint() + return TmxHydroExitSavepoint( + savepoint, self.serializer, size=self.grid_size, backend=self.backend + ) + + def from_savepoint_tmx_temperature_exit(self, date: str) -> TmxTemperatureExitSavepoint: + savepoint = ( + self.serializer.savepoint["tmx-temperature-exit"].id[1].date[date].as_savepoint() + ) + return TmxTemperatureExitSavepoint( + savepoint, self.serializer, size=self.grid_size, backend=self.backend + ) From 8715a4b921bb3700d1e0d55339ff04d84fe3a1da Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 20:27:26 +0200 Subject: [PATCH 09/18] Read the tmx time step from the experiment's namelist in the scalar-diffusion datatests --- .../tmx/tests/tmx/fixtures.py | 18 +++++++++++++++++- .../test_tmx_scalar_diffusion.py | 7 ++++--- .../tmx/tests/tmx/integration_tests/utils.py | 3 --- 3 files changed, 21 insertions(+), 7 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py index 8046f5db32..b6bb155fd0 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py @@ -10,7 +10,7 @@ from icon4py.model.atmosphere.subgrid_scale_physics.tmx.config import TmxConfig from icon4py.model.common.decomposition import definitions as decomposition -from icon4py.model.common.utils import fortran_config +from icon4py.model.common.utils import fortran_config, time_utils from icon4py.model.testing import datatest_utils as dt_utils, definitions from icon4py.model.testing.fixtures.datatest import ( backend, @@ -39,3 +39,19 @@ def tmx_config( fname=fortran_config.ATM_DICT_FNAME, ) return TmxConfig.from_fortran_dict(atm_dict=atm_dict) + + +@pytest.fixture +def tmx_dtime( + experiment_description: definitions.ExperimentDescription, + process_props: decomposition.ProcessProperties, + download_ser_data: None, # downloads data as side-effect +) -> float: + """The tmx time step [s]: ``dt_vdf`` of the experiment's ``aes_phy_nml``.""" + input_dict = dt_utils.load_fortran_dict( + experiment_description=experiment_description, + process_props=process_props, + fname=fortran_config.INPUT_DICT_FNAME, + ) + dt_vdf = input_dict["aes_phy_nml"]["aes_phy_config"][0]["dt_vdf"] + return time_utils.relativetime_from_iso8601(dt_vdf).total_seconds() diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py index 1ed1935755..fb7a222997 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py @@ -30,7 +30,6 @@ from .utils import ( RTOL, TMX_DATES, - TMX_DTIME, construct_input_state, construct_interpolation_state, construct_metric_state, @@ -115,6 +114,7 @@ def test_tmx_run_hydrometeor_diffusion_single_step( backend: gtx_typing.Backend | None, date: str, tmx_config: TmxConfig, + tmx_dtime: float, ) -> None: setup = _setup( data_provider=data_provider, @@ -134,7 +134,7 @@ def test_tmx_run_hydrometeor_diffusion_single_step( diagnostic_state=setup.diagnostic_state, tendency_state=setup.tendency_state, new_state=setup.new_state, - dtime=TMX_DTIME, + dtime=tmx_dtime, ) fields = ( @@ -166,6 +166,7 @@ def test_tmx_run_temperature_diffusion_single_step( backend: gtx_typing.Backend | None, date: str, tmx_config: TmxConfig, + tmx_dtime: float, ) -> None: setup = _setup( data_provider=data_provider, @@ -192,7 +193,7 @@ def test_tmx_run_temperature_diffusion_single_step( diagnostic_state=setup.diagnostic_state, tendency_state=setup.tendency_state, new_state=new_state, - dtime=TMX_DTIME, + dtime=tmx_dtime, ) fields = ( diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py index 1747246e18..c9761c15eb 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py @@ -32,9 +32,6 @@ # so the verification tests parametrize over the subsequent steps only. TMX_DATES: tuple[str, ...] = ("2008-09-01T00:05:00.000", "2008-09-01T00:10:00.000") -# 'dt_vdf' of the archive's aes_phy_config [s]. -TMX_DTIME: float = 300.0 - # Relative tolerance of all tmx integration datatests. RTOL: float = 3.0e-12 From 4f965564c9a22043c41ce02f4140a7619909de53 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 21:10:27 +0200 Subject: [PATCH 10/18] Reject the explicit solver in the tmx scalar diffusion --- .../atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py | 5 +++++ .../tests/tmx/integration_tests/test_tmx_scalar_diffusion.py | 1 + 2 files changed, 6 insertions(+) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py index e2da1d09a5..bdf252d83a 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py @@ -55,7 +55,12 @@ def __init__( energy_type: tmx_config.EnergyType, use_scale_turb_energy_flux: bool, scale_turb_energy_flux: float, + solver_type: tmx_config.SolverType, ) -> None: + if solver_type != tmx_config.SolverType.IMPLICIT: + raise NotImplementedError( + "the scalar diffusion only implements the implicit vertical diffusion solver." + ) self._exchange = exchange # ``zfactor`` in Compute_diffusion_temperature energy_flux_factor = scale_turb_energy_flux if use_scale_turb_energy_flux else 1.0 diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py index fb7a222997..3294486aa2 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py @@ -82,6 +82,7 @@ def _setup( energy_type=config.energy_type, use_scale_turb_energy_flux=config.use_scale_turb_energy_flux, scale_turb_energy_flux=config.scale_turb_energy_flux, + solver_type=config.solver_type, ) return _Setup( component=component, From 29f8f8e74106b9c212be7663b60c76bf43517338 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 21:10:27 +0200 Subject: [PATCH 11/18] Test the tmx scalar-diffusion programs on a horizontal domain inside the field --- .../tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py index b301a6bff1..5eaeca7b7f 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py @@ -120,9 +120,10 @@ def diffuse_scalar_numpy( def _cells(grid: base.Grid) -> tuple[gtx.int32, gtx.int32]: cell_domain = h_grid.domain(dims.CellDim) + # strictly inside the field, so that an output written on the whole field is caught return ( - grid.start_index(cell_domain(h_grid.Zone.NUDGING)), - grid.end_index(cell_domain(h_grid.Zone.LOCAL)), + grid.start_index(cell_domain(h_grid.Zone.NUDGING)) + 1, + grid.end_index(cell_domain(h_grid.Zone.LOCAL)) - 1, ) From d9c5dc31e622d88d0a810d526eade1afc1248149 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 21:16:39 +0200 Subject: [PATCH 12/18] Keep only the surface fluxes the tmx scalar diffusion reads --- .../subgrid_scale_physics/tmx/tmx_states.py | 25 ------------------- .../tmx/tests/tmx/integration_tests/utils.py | 3 --- .../src/icon4py/model/testing/serialbox.py | 9 ------- 3 files changed, 37 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py index 713a9bd753..73135fa704 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py @@ -90,31 +90,6 @@ class TmxSurfaceFluxState: """Surface evapotranspiration flux (``evspsbl``) [kg/(m^2 s)].""" sensible_heat_flux: fa.CellField[ta.wpfloat] """Surface sensible heat flux (``hfss``) [W/m^2].""" - u_stress: fa.CellField[ta.wpfloat] - """Zonal surface wind stress (``tauu``) [N/m^2].""" - v_stress: fa.CellField[ta.wpfloat] - """Meridional surface wind stress (``tauv``) [N/m^2].""" - q_snocpymlt: fa.CellField[ta.wpfloat] - """Heating used to melt snow on the canopy [W/m^2].""" - - @classmethod - def allocate( - cls, grid: base_grid.Grid, allocator: gtx_typing.Allocator | None = None - ) -> TmxSurfaceFluxState: - """Allocate a surface flux state with all fields initialized to zero.""" - - def surface(horizontal_dim: gtx.Dimension) -> gtx.Field: - return data_alloc.zero_field( - grid, horizontal_dim, dtype=ta.wpfloat, allocator=allocator - ) - - return cls( - evapotranspiration=surface(dims.CellDim), - sensible_heat_flux=surface(dims.CellDim), - u_stress=surface(dims.CellDim), - v_stress=surface(dims.CellDim), - q_snocpymlt=surface(dims.CellDim), - ) @dataclasses.dataclass(frozen=True) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py index c9761c15eb..6c5ebc3a8a 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py @@ -113,7 +113,4 @@ def construct_surface_flux_state( return tmx_states.TmxSurfaceFluxState( evapotranspiration=surface_fluxes_savepoint.evspsbl(), sensible_heat_flux=surface_fluxes_savepoint.hfss(), - u_stress=surface_fluxes_savepoint.tauu(), - v_stress=surface_fluxes_savepoint.tauv(), - q_snocpymlt=surface_fluxes_savepoint.q_snocpymlt(), ) diff --git a/model/testing/src/icon4py/model/testing/serialbox.py b/model/testing/src/icon4py/model/testing/serialbox.py index 398165b3e8..8cb8b30ea6 100644 --- a/model/testing/src/icon4py/model/testing/serialbox.py +++ b/model/testing/src/icon4py/model/testing/serialbox.py @@ -2060,15 +2060,6 @@ def evspsbl(self): def hfss(self): return self._get_field("hfss", dims.CellDim) - def tauu(self): - return self._get_field("tauu", dims.CellDim) - - def tauv(self): - return self._get_field("tauv", dims.CellDim) - - def q_snocpymlt(self): - return self._get_field("q_snocpymlt", dims.CellDim) - class TmxDiagnosticsExitSavepoint(IconSavepoint): """Savepoint at exit of vdf Compute_diagnostics in mo_vdf_atmo.f90.""" From 9d5548459f6b09fe8828baf5bae89f8a365a251d Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Wed, 23 Sep 2026 21:17:51 +0200 Subject: [PATCH 13/18] Keep the surface-flux savepoint and constructor where the wind diffusion adds them --- .../tmx/tests/tmx/integration_tests/utils.py | 2 +- model/testing/src/icon4py/model/testing/serialbox.py | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py index 6c5ebc3a8a..60fde7c82d 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py @@ -102,8 +102,8 @@ def construct_input_state(entry_savepoint: sb.TmxEntrySavepoint) -> tmx_states.T qr=entry_savepoint.qr(), qs=entry_savepoint.qs(), qg=entry_savepoint.qg(), - rho=entry_savepoint.rho(), air_mass=entry_savepoint.mair(), + rho=entry_savepoint.rho(), ) diff --git a/model/testing/src/icon4py/model/testing/serialbox.py b/model/testing/src/icon4py/model/testing/serialbox.py index 8cb8b30ea6..e432905940 100644 --- a/model/testing/src/icon4py/model/testing/serialbox.py +++ b/model/testing/src/icon4py/model/testing/serialbox.py @@ -2041,15 +2041,15 @@ def qg(self): def rho(self): return self._get_field("rho", dims.CellDim, dims.KDim) + def mair(self): + return self._get_field("mair", dims.CellDim, dims.KDim) + def tempv(self): return self._get_field("tempv", dims.CellDim, dims.KDim) def pres(self): return self._get_field("pres", dims.CellDim, dims.KDim) - def mair(self): - return self._get_field("mair", dims.CellDim, dims.KDim) - class TmxSurfaceFluxesSavepoint(IconSavepoint): """Savepoint after the surface model call in vdf Compute in mo_vdf.f90.""" From 0842d78593a1b32f43af95767ae800766442d704 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 16:35:37 +0200 Subject: [PATCH 14/18] Act on the #1490 review: separate the energy types and drop the explicit solver The energy type decides three operations (energy from temperature, surface energy flux, temperature from energy); each is now a named pair of field operators selected by the static use_internal_energy, and the energy goes through the same diffusion core as the tracers. The energy matrix is assembled inside the energy program. SolverType has only IMPLICIT; TmxConfig and namelist parsing reject the explicit solver, which was only used for early testing in ICON. --- .../subgrid_scale_physics/tmx/config.py | 21 +- .../tmx/scalar_diffusion.py | 37 +--- .../tmx/stencils/scalar_diffusion.py | 182 +++++++++++------- .../test_tmx_scalar_diffusion.py | 1 - .../stencil_tests/test_scalar_diffusion.py | 31 ++- .../tests/tmx/unit_tests/test_tmx_config.py | 22 ++- 6 files changed, 176 insertions(+), 118 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py index 0598193aa5..4ab28001b3 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py @@ -24,9 +24,8 @@ @config_io.register_enum class SolverType(int, enum.Enum): - """Type of the vertical diffusion solver.""" + """Type of the vertical diffusion solver; ICON's explicit solver (1) is not ported.""" - EXPLICIT = 1 # explicit time stepping IMPLICIT = 2 # implicit time stepping @@ -48,7 +47,7 @@ class TmxConfig: solver_type: typing.Annotated[ SolverType, common_conf_opt.ConfigOption( - description="Type of the vertical diffusion solver (explicit or implicit).", + description="Type of the vertical diffusion solver (only implicit).", icon_equivalent=common_conf_opt.IconOption( "solver_type", ("aes_vdf_nml", "aes_vdf_config"), unnamed_index=23 ), @@ -201,7 +200,13 @@ class TmxConfig: ] = 300.0 def __post_init__(self) -> None: - self.solver_type = SolverType(self.solver_type) + try: + self.solver_type = SolverType(self.solver_type) + except ValueError: + raise ValueError( + f"Invalid argument 'solver_type': only the implicit solver " + f"({SolverType.IMPLICIT.value}) is implemented, got {self.solver_type}." + ) from None self.energy_type = EnergyType(self.energy_type) if self.turb_prandtl <= 0.0: @@ -226,6 +231,7 @@ def from_fortran_dict(cls, *, atm_dict: dict[str, Any], **overrides: Any) -> Tmx # Keep these values and the options' unnamed_index positions in sync num_members = 42 use_tmx_index = 22 + solver_type_index = 23 flat = atm_dict["aes_vdf_nml"]["aes_vdf_config"] if len(flat) % num_members != 0: @@ -241,4 +247,11 @@ def from_fortran_dict(cls, *, atm_dict: dict[str, Any], **overrides: Any) -> Tmx f"'aes_vdf_config', found {use_tmx!r}: either the run does not use tmx or " "the t_vdiff_config member order changed." ) + # checked here because the enum conversion of the option fails before __post_init__ + solver_type = flat[solver_type_index] + if solver_type != SolverType.IMPLICIT.value: + raise ValueError( + f"Invalid 'solver_type' {solver_type!r} in 'aes_vdf_config': only the implicit " + f"solver ({SolverType.IMPLICIT.value}) is implemented." + ) return common_conf_opt.construct_config_from_icon(cls, atm_dict, **overrides) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py index bdf252d83a..9ec15468e0 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py @@ -55,12 +55,7 @@ def __init__( energy_type: tmx_config.EnergyType, use_scale_turb_energy_flux: bool, scale_turb_energy_flux: float, - solver_type: tmx_config.SolverType, ) -> None: - if solver_type != tmx_config.SolverType.IMPLICIT: - raise NotImplementedError( - "the scalar diffusion only implements the implicit vertical diffusion solver." - ) self._exchange = exchange # ``zfactor`` in Compute_diffusion_temperature energy_flux_factor = scale_turb_energy_flux if use_scale_turb_energy_flux else 1.0 @@ -73,7 +68,7 @@ def __init__( zero_field = functools.partial( data_alloc.zero_field, grid, allocator=model_backends.get_allocator(backend) ) - # the matrix is assembled once for the three tracers and once for the energy + # assembled once and reused for the three tracers self._matrix_a: fa.CellKField[ta.wpfloat] = zero_field(dims.CellDim, dims.KDim) self._matrix_b: fa.CellKField[ta.wpfloat] = zero_field(dims.CellDim, dims.KDim) self._matrix_c: fa.CellKField[ta.wpfloat] = zero_field(dims.CellDim, dims.KDim) @@ -88,23 +83,14 @@ def __init__( "vertical_start": gtx.int32(0), "vertical_end": gtx.int32(grid.num_levels), } - assemble_matrix = functools.partial( - setup_program, + self.assemble_tracer_diffusion_matrix = setup_program( backend=backend, program=scalar_stencils.assemble_scalar_diffusion_matrix, + constant_args={"inv_dz": metric_state.inv_ddqz_z_half, "prefactor": 1.0}, horizontal_sizes=horizontal_sizes, vertical_sizes=vertical_sizes, offset_provider={}, ) - self.assemble_tracer_diffusion_matrix = assemble_matrix( - constant_args={"inv_dz": metric_state.inv_ddqz_z_half, "prefactor": 1.0} - ) - self.assemble_energy_diffusion_matrix = assemble_matrix( - constant_args={ - "inv_dz": metric_state.inv_ddqz_z_half, - "prefactor": energy_flux_factor, - } - ) horizontal_diffusion_args = { "inv_dual_edge_length": edge_params.inverse_dual_edge_lengths, "geofac_div": interpolation_state.geofac_div, @@ -135,6 +121,7 @@ def __init__( program=scalar_stencils.diffuse_energy_and_update_temperature, constant_args={ **horizontal_diffusion_args, + "inv_dz": metric_state.inv_ddqz_z_half, "height_above_ground": metric_state.height_above_ground, "prefactor": energy_flux_factor, "grav": constants.GRAV, @@ -237,24 +224,12 @@ def run_temperature_diffusion( ) log.debug("communication of energy (cells): start") - energy_exchange = self._exchange.start(dims.CellDim, self.energy) - - self.assemble_energy_diffusion_matrix( - diffusivity=diagnostic_state.kh_ic, - air_mass=input_state.air_mass, - a=self._matrix_a, - b=self._matrix_b, - c=self._matrix_c, - ) - - energy_exchange.finish() + self._exchange.exchange(dims.CellDim, self.energy) log.debug("communication of energy (cells): end") self.diffuse_energy_and_update_temperature( energy=self.energy, - a=self._matrix_a, - b=self._matrix_b, - c=self._matrix_c, + diffusivity=diagnostic_state.kh_ic, sensible_heat_flux=surface_flux_state.sensible_heat_flux, evapotranspiration=surface_flux_state.evapotranspiration, temperature=input_state.temperature, diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py index 90b2658db2..55004f4679 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py @@ -76,7 +76,8 @@ def _diffuse_scalar( a: fa.CellKField[wpfloat], b: fa.CellKField[wpfloat], c: fa.CellKField[wpfloat], - rhs: fa.CellKField[wpfloat], + surface_flux: fa.CellKField[wpfloat], + air_mass: fa.CellKField[wpfloat], rho: fa.CellKField[wpfloat], km_ie: fa.EdgeKHalfField[wpfloat], inv_dual_edge_length: fa.EdgeField[wpfloat], @@ -84,14 +85,19 @@ def _diffuse_scalar( rturb_prandtl: wpfloat, prefactor: wpfloat, dtime: wpfloat, + maxlvl: gtx.int32, ) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: """ New value and tendency of a cell scalar after one diffusion step: implicit in the vertical - with matrix (a, b, c), explicit and conservative in the horizontal. + with matrix (a, b, c) and the surface flux entering the bottom row maxlvl, explicit and + conservative in the horizontal. - var must be valid on the halo cells. + Only the bottom row of surface_flux is read. var must be valid on the halo cells. """ zero = wpfloat("0.0") * var + rhs = concat_where( + dims.KDim < maxlvl, zero, -surface_flux * prefactor * (wpfloat("1.0") / air_mass) + ) vertical_tend = _solve_implicit_vertical_diffusion_on_cells(var, a, b, c, rhs, zero, dtime) flux = ( wpfloat("0.5") @@ -122,18 +128,13 @@ def _diffuse_tracer( dtime: wpfloat, maxlvl: gtx.int32, ) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: - """:func:`_diffuse_scalar` with the surface flux entering the bottom row maxlvl.""" - rhs = concat_where( - dims.KDim < maxlvl, - wpfloat("0.0") * var, - wpfloat("0.0") - surface_flux * prefactor * (wpfloat("1.0") / air_mass), - ) return _diffuse_scalar( var, a, b, c, - rhs, + broadcast(surface_flux, (dims.CellDim, dims.KDim)), + air_mass, rho, km_ie, inv_dual_edge_length, @@ -141,6 +142,7 @@ def _diffuse_tracer( rturb_prandtl, prefactor, dtime, + maxlvl, ) @@ -189,6 +191,71 @@ def diffuse_tracer( ) +@gtx.field_operator +def _compute_internal_energy_from_temperature( + temperature: fa.CellKField[wpfloat], + qv: fa.CellKField[wpfloat], + q_liquid: fa.CellKField[wpfloat], + q_solid: fa.CellKField[wpfloat], + height_above_ground: fa.CellKField[wpfloat], + grav: wpfloat, +) -> fa.CellKField[wpfloat]: + """Specific internal energy plus cvd / cpd times the geopotential above ground.""" + one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) + return ( + compute_internal_energy_per_area(temperature, qv, q_liquid, q_solid, one, one) + + grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd + ) + + +@gtx.field_operator +def _compute_temperature_from_internal_energy( + energy: fa.CellKField[wpfloat], + qv: fa.CellKField[wpfloat], + q_liquid: fa.CellKField[wpfloat], + q_solid: fa.CellKField[wpfloat], + height_above_ground: fa.CellKField[wpfloat], + grav: wpfloat, +) -> fa.CellKField[wpfloat]: + """Inverse of :func:`_compute_internal_energy_from_temperature`.""" + one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) + return compute_temperature_from_internal_energy_per_area( + energy - grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd, + qv, + q_liquid, + q_solid, + one, + one, + ) + + +@gtx.field_operator +def _compute_temperature_from_dry_static_energy( + energy: fa.CellKField[wpfloat], + height_above_ground: fa.CellKField[wpfloat], + grav: wpfloat, +) -> fa.CellKField[wpfloat]: + return (energy - grav * height_above_ground) / PhysicsConstants.cpd + + +@gtx.field_operator +def _compute_surface_internal_energy_flux( + sensible_heat_flux: fa.CellField[wpfloat], + evapotranspiration: fa.CellField[wpfloat], + temperature: fa.CellKField[wpfloat], +) -> fa.CellKField[wpfloat]: + return sensible_heat_flux + temperature * evapotranspiration * ( + PhysicsConstants.cvv - PhysicsConstants.cvd + ) + + +@gtx.field_operator +def _compute_surface_dry_static_energy_flux( + sensible_heat_flux: fa.CellField[wpfloat], +) -> fa.CellField[wpfloat]: + return sensible_heat_flux * PhysicsConstants.cpd / PhysicsConstants.cvd + + @gtx.field_operator def _compute_energy_from_temperature( temperature: fa.CellKField[wpfloat], @@ -202,19 +269,13 @@ def _compute_energy_from_temperature( grav: wpfloat, use_internal_energy: bool, ) -> fa.CellKField[wpfloat]: - """ - Specific energy diffused by the heat diffusion: the internal energy plus cvd / cpd times the - geopotential above ground, or the dry static energy. - """ - if use_internal_energy: - one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) - energy = ( - compute_internal_energy_per_area(temperature, qv, qc + qr, qi + qs + qg, one, one) - + grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd + return ( + _compute_internal_energy_from_temperature( + temperature, qv, qc + qr, qi + qs + qg, height_above_ground, grav ) - else: - energy = _compute_dry_static_energy(temperature, height_above_ground, grav) - return energy + if use_internal_energy + else _compute_dry_static_energy(temperature, height_above_ground, grav) + ) @gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) @@ -257,9 +318,8 @@ def compute_energy_from_temperature( @gtx.field_operator def _diffuse_energy_and_update_temperature( energy: fa.CellKField[wpfloat], - a: fa.CellKField[wpfloat], - b: fa.CellKField[wpfloat], - c: fa.CellKField[wpfloat], + diffusivity: fa.CellKHalfField[wpfloat], + inv_dz: fa.CellKHalfField[wpfloat], sensible_heat_flux: fa.CellField[wpfloat], evapotranspiration: fa.CellField[wpfloat], temperature: fa.CellKField[wpfloat], @@ -279,45 +339,32 @@ def _diffuse_energy_and_update_temperature( prefactor: wpfloat, grav: wpfloat, dtime: wpfloat, + minlvl: gtx.int32, maxlvl: gtx.int32, use_internal_energy: bool, ) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: """ New temperature and its tendency after one diffusion step of the energy of - :func:`_compute_energy_from_temperature`, with the surface energy flux entering the bottom - row maxlvl. - - The new temperature is recovered with the new qv, qc and qi. + :func:`_compute_energy_from_temperature`, converted back with the new qv, qc and qi. """ - inv_air_mass = wpfloat("1.0") / air_mass - # only the bottom row of the surface term is used, where 'temperature' is that of the - # lowest level - if use_internal_energy: - surface_term = ( - wpfloat("0.0") - - ( - sensible_heat_flux - + temperature * evapotranspiration * (PhysicsConstants.cvv - PhysicsConstants.cvd) - ) - * prefactor - * inv_air_mass - ) - else: - surface_term = ( - wpfloat("0.0") - - sensible_heat_flux - * PhysicsConstants.cpd - / PhysicsConstants.cvd - * prefactor - * inv_air_mass + a, b, c = _assemble_scalar_diffusion_matrix( + diffusivity, inv_dz, air_mass, prefactor, minlvl, maxlvl + ) + # only the bottom row is used, where 'temperature' is that of the lowest level + surface_flux = ( + _compute_surface_internal_energy_flux(sensible_heat_flux, evapotranspiration, temperature) + if use_internal_energy + else broadcast( + _compute_surface_dry_static_energy_flux(sensible_heat_flux), (dims.CellDim, dims.KDim) ) - rhs = concat_where(dims.KDim < maxlvl, wpfloat("0.0") * energy, surface_term) + ) new_energy, _ = _diffuse_scalar( energy, a, b, c, - rhs, + surface_flux, + air_mass, rho, km_ie, inv_dual_edge_length, @@ -325,28 +372,25 @@ def _diffuse_energy_and_update_temperature( rturb_prandtl, prefactor, dtime, + maxlvl, ) - if use_internal_energy: - one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) - new_temperature = compute_temperature_from_internal_energy_per_area( - new_energy - grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd, - new_qv, - new_qc + qr, - new_qi + qs + qg, - one, - one, + q_liquid = new_qc + qr + q_solid = new_qi + qs + qg + new_temperature = ( + _compute_temperature_from_internal_energy( + new_energy, new_qv, q_liquid, q_solid, height_above_ground, grav ) - else: - new_temperature = (new_energy - grav * height_above_ground) / PhysicsConstants.cpd + if use_internal_energy + else _compute_temperature_from_dry_static_energy(new_energy, height_above_ground, grav) + ) return new_temperature, (new_temperature - temperature) * (wpfloat("1.0") / dtime) @gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) def diffuse_energy_and_update_temperature( energy: fa.CellKField[wpfloat], - a: fa.CellKField[wpfloat], - b: fa.CellKField[wpfloat], - c: fa.CellKField[wpfloat], + diffusivity: fa.CellKHalfField[wpfloat], + inv_dz: fa.CellKHalfField[wpfloat], sensible_heat_flux: fa.CellField[wpfloat], evapotranspiration: fa.CellField[wpfloat], temperature: fa.CellKField[wpfloat], @@ -376,9 +420,8 @@ def diffuse_energy_and_update_temperature( ) -> None: _diffuse_energy_and_update_temperature( energy, - a, - b, - c, + diffusivity, + inv_dz, sensible_heat_flux, evapotranspiration, temperature, @@ -398,6 +441,7 @@ def diffuse_energy_and_update_temperature( prefactor, grav, dtime, + vertical_start, vertical_end - 1, use_internal_energy, out=(new_temperature, tend_temperature), diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py index 3294486aa2..fb7a222997 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py @@ -82,7 +82,6 @@ def _setup( energy_type=config.energy_type, use_scale_turb_energy_flux=config.use_scale_turb_energy_flux, scale_turb_energy_flux=config.scale_turb_energy_flux, - solver_type=config.solver_type, ) return _Setup( component=component, diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py index 5eaeca7b7f..6d60574dff 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py @@ -23,7 +23,11 @@ from icon4py.model.common.type_alias import wpfloat from icon4py.model.testing import stencil_tests -from .test_vertical_diffusion import implicit_diffusion_tendency_numpy +from .test_vertical_diffusion import ( + diffusion_matrix_numpy, + implicit_diffusion_tendency_numpy, + matrix_diagonals_on_rows, +) _DOMAIN_ARGS = ("horizontal_start", "horizontal_end", "vertical_start", "vertical_end") @@ -133,9 +137,6 @@ def _diffusion_input_data( horizontal_start, horizontal_end = _cells(grid) return { var_name: data_alloc.random_field(dims.CellDim, dims.KDim), - "a": data_alloc.random_field(dims.CellDim, dims.KDim, low=-1.0, high=0.0), - "b": data_alloc.random_field(dims.CellDim, dims.KDim, low=2.0, high=3.0), - "c": data_alloc.random_field(dims.CellDim, dims.KDim, low=-1.0, high=0.0), "air_mass": data_alloc.random_field(dims.CellDim, dims.KDim, low=1.0, high=2.0), "rho": data_alloc.random_field(dims.CellDim, dims.KDim, low=0.5, high=1.5), "km_ie": data_alloc.random_field(dims.EdgeDim, dims.KHalfDim, low=0.0), @@ -214,6 +215,9 @@ def reference( def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: return dict( **_diffusion_input_data(data_alloc, grid, "var"), + a=data_alloc.random_field(dims.CellDim, dims.KDim, low=-1.0, high=0.0), + b=data_alloc.random_field(dims.CellDim, dims.KDim, low=2.0, high=3.0), + c=data_alloc.random_field(dims.CellDim, dims.KDim, low=-1.0, high=0.0), surface_flux=data_alloc.random_field(dims.CellDim), new_var=data_alloc.zero_field(dims.CellDim, dims.KDim), tend=data_alloc.zero_field(dims.CellDim, dims.KDim), @@ -345,6 +349,9 @@ def reference( qs: np.ndarray, qg: np.ndarray, height_above_ground: np.ndarray, + diffusivity: np.ndarray, + inv_dz: np.ndarray, + air_mass: np.ndarray, prefactor: float, grav: float, dtime: float, @@ -356,6 +363,12 @@ def reference( **kwargs: Any, ) -> dict: cells, rows = _slices(horizontal_start, horizontal_end, vertical_start, vertical_end) + interfaces = slice(rows.start + 1, rows.stop) + matrix = diffusion_matrix_numpy( + prefactor * diffusivity[:, interfaces] * inv_dz[:, interfaces], + 1.0 / air_mass[:, rows], + ) + a, b, c = matrix_diagonals_on_rows(matrix, air_mass.shape, rows) if use_internal_energy: temperature_sfc = temperature[:, vertical_end - 1] surface_flux = sensible_heat_flux + temperature_sfc * evapotranspiration * ( @@ -366,6 +379,10 @@ def reference( new_energy, _ = diffuse_scalar_numpy( stencil_tests.connectivities_asnumpy(grid), var=energy, + a=a, + b=b, + c=c, + air_mass=air_mass, surface_flux=surface_flux, prefactor=prefactor, dtime=dtime, @@ -374,10 +391,6 @@ def reference( **{ name: kwargs[name] for name in ( - "a", - "b", - "c", - "air_mass", "rho", "km_ie", "inv_dual_edge_length", @@ -409,6 +422,8 @@ def _diffuse_energy_input_data( ) -> dict: return dict( **_diffusion_input_data(data_alloc, grid, "energy"), + diffusivity=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.0), + inv_dz=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.1), **_tracers(data_alloc, prefix="new_"), sensible_heat_flux=data_alloc.random_field(dims.CellDim), evapotranspiration=data_alloc.random_field(dims.CellDim), diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_tmx_config.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_tmx_config.py index ffa1f52da9..e31ef20201 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_tmx_config.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_tmx_config.py @@ -25,14 +25,20 @@ def test_config_rejects_negative_km_min() -> None: def test_config_coerces_enums_from_ints() -> None: - config = tmx_config.TmxConfig(solver_type=1, energy_type=1) - assert config.solver_type is tmx_config.SolverType.EXPLICIT + config = tmx_config.TmxConfig(solver_type=2, energy_type=1) + assert config.solver_type is tmx_config.SolverType.IMPLICIT assert config.energy_type is tmx_config.EnergyType.DRY_STATIC +@pytest.mark.parametrize("solver_type", [1, 3]) +def test_config_rejects_unimplemented_solver_types(solver_type: int) -> None: + with pytest.raises(ValueError, match="solver_type"): + tmx_config.TmxConfig(solver_type=solver_type) + + def test_config_rejects_invalid_enum_values() -> None: with pytest.raises(ValueError): - tmx_config.TmxConfig(solver_type=3) + tmx_config.TmxConfig(energy_type=3) def _echoed_vdf_record(**overrides: object) -> list[object]: @@ -70,7 +76,7 @@ def test_config_from_fortran_dict() -> None: fortran_dict = { "aes_vdf_nml": { "aes_vdf_config": _echoed_vdf_record( - solver_type=1, + solver_type=2, energy_type=1, dissipation_factor=0.5, use_louis=False, @@ -89,7 +95,7 @@ def test_config_from_fortran_dict() -> None: } } config = tmx_config.TmxConfig.from_fortran_dict(atm_dict=fortran_dict) - assert config.solver_type is tmx_config.SolverType.EXPLICIT + assert config.solver_type is tmx_config.SolverType.IMPLICIT assert config.energy_type is tmx_config.EnergyType.DRY_STATIC assert config.dissipation_factor == 0.5 assert config.use_louis is False @@ -106,6 +112,12 @@ def test_config_from_fortran_dict() -> None: assert config.max_turb_scale == 150.0 +def test_config_from_fortran_dict_rejects_explicit_solver() -> None: + record = _echoed_vdf_record(solver_type=1) + with pytest.raises(ValueError, match="only the implicit solver"): + tmx_config.TmxConfig.from_fortran_dict(atm_dict={"aes_vdf_nml": {"aes_vdf_config": record}}) + + def test_config_from_fortran_dict_rejects_changed_member_count() -> None: record = _echoed_vdf_record() with pytest.raises(ValueError, match="not a multiple"): From 44f879c16b5a64c09a4746a4b5cb8989f931144f Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Fri, 25 Sep 2026 18:15:51 +0200 Subject: [PATCH 15/18] Parametrize the scalar-diffusion energy stencil tests over the energy type --- .../stencil_tests/test_scalar_diffusion.py | 57 ++++++++----------- 1 file changed, 23 insertions(+), 34 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py index 6d60574dff..a7f0594e34 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py @@ -11,6 +11,7 @@ import gt4py.next as gtx import numpy as np +import pytest from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils.scalar_diffusion import ( compute_energy_from_temperature, @@ -239,7 +240,7 @@ def _tracers(data_alloc: stencil_tests.DataAllocationWrapper, prefix: str = "") } -class _ComputeEnergyFromTemperature: +class TestComputeEnergyFromTemperature(stencil_tests.StencilTest): PROGRAM = compute_energy_from_temperature OUTPUTS = ("energy",) STATIC_PARAMS = { @@ -285,6 +286,16 @@ def reference( )[cells, rows] return dict(energy=energy) + @stencil_tests.input_data_fixture( + params=[True, False], ids=["internal_energy", "dry_static_energy"] + ) + def input_data( + data_alloc: stencil_tests.DataAllocationWrapper, + grid: base.Grid, + request: pytest.FixtureRequest, + ) -> dict: + return _compute_energy_input_data(data_alloc, grid, use_internal_energy=request.param) + def _compute_energy_input_data( data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, use_internal_energy: bool @@ -304,23 +315,7 @@ def _compute_energy_input_data( ) -class TestComputeInternalEnergyFromTemperature( - _ComputeEnergyFromTemperature, stencil_tests.StencilTest -): - @stencil_tests.input_data_fixture - def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: - return _compute_energy_input_data(data_alloc, grid, use_internal_energy=True) - - -class TestComputeDryStaticEnergyFromTemperature( - _ComputeEnergyFromTemperature, stencil_tests.StencilTest -): - @stencil_tests.input_data_fixture - def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: - return _compute_energy_input_data(data_alloc, grid, use_internal_energy=False) - - -class _DiffuseEnergyAndUpdateTemperature: +class TestDiffuseEnergyAndUpdateTemperature(stencil_tests.StencilTest): PROGRAM = diffuse_energy_and_update_temperature OUTPUTS = ("new_temperature", "tend_temperature") STATIC_PARAMS = { @@ -416,6 +411,16 @@ def reference( ) / dtime return dict(new_temperature=new_temperature, tend_temperature=tend_temperature) + @stencil_tests.input_data_fixture( + params=[True, False], ids=["internal_energy", "dry_static_energy"] + ) + def input_data( + data_alloc: stencil_tests.DataAllocationWrapper, + grid: base.Grid, + request: pytest.FixtureRequest, + ) -> dict: + return _diffuse_energy_input_data(data_alloc, grid, use_internal_energy=request.param) + def _diffuse_energy_input_data( data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, use_internal_energy: bool @@ -435,19 +440,3 @@ def _diffuse_energy_input_data( grav=constants.GRAV, use_internal_energy=use_internal_energy, ) - - -class TestDiffuseInternalEnergyAndUpdateTemperature( - _DiffuseEnergyAndUpdateTemperature, stencil_tests.StencilTest -): - @stencil_tests.input_data_fixture - def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: - return _diffuse_energy_input_data(data_alloc, grid, use_internal_energy=True) - - -class TestDiffuseDryStaticEnergyAndUpdateTemperature( - _DiffuseEnergyAndUpdateTemperature, stencil_tests.StencilTest -): - @stencil_tests.input_data_fixture - def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: - return _diffuse_energy_input_data(data_alloc, grid, use_internal_energy=False) From 5c856d575ecade14fd9b8dcff8837405de5ed6d3 Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Mon, 28 Sep 2026 13:53:50 +0200 Subject: [PATCH 16/18] docs: act on the second #1490 review round --- .../subgrid_scale_physics/tmx/scalar_diffusion.py | 2 +- .../subgrid_scale_physics/tmx/stencils/scalar_diffusion.py | 7 ++++--- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py index 9ec15468e0..1ce33f9aeb 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py @@ -207,7 +207,7 @@ def run_temperature_diffusion( in mo_vdf.f90). The energy is computed with the input tracers and converted back with the new qv, qc and - qi of ``new_state``, so this runs after :meth:`run_hydrometeor_diffusion`. Needs + qi of ``new_state``, so this runs after ``run_hydrometeor_diffusion``. Needs ``kh_ic`` and ``km_ie`` of ``diagnostic_state``. """ log.debug("tmx temperature diffusion (Compute_diffusion_temperature): start") diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py index 55004f4679..c80379707f 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py @@ -128,12 +128,13 @@ def _diffuse_tracer( dtime: wpfloat, maxlvl: gtx.int32, ) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: + surface_flux_on_levels = broadcast(surface_flux, (dims.CellDim, dims.KDim)) return _diffuse_scalar( var, a, b, c, - broadcast(surface_flux, (dims.CellDim, dims.KDim)), + surface_flux_on_levels, air_mass, rho, km_ie, @@ -217,7 +218,7 @@ def _compute_temperature_from_internal_energy( height_above_ground: fa.CellKField[wpfloat], grav: wpfloat, ) -> fa.CellKField[wpfloat]: - """Inverse of :func:`_compute_internal_energy_from_temperature`.""" + """Inverse of ``_compute_internal_energy_from_temperature``.""" one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) return compute_temperature_from_internal_energy_per_area( energy - grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd, @@ -345,7 +346,7 @@ def _diffuse_energy_and_update_temperature( ) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: """ New temperature and its tendency after one diffusion step of the energy of - :func:`_compute_energy_from_temperature`, converted back with the new qv, qc and qi. + ``_compute_energy_from_temperature``, converted back with the new qv, qc and qi. """ a, b, c = _assemble_scalar_diffusion_matrix( diffusivity, inv_dz, air_mass, prefactor, minlvl, maxlvl From 7d018ec5fd62b9b62b7df1fa08ec000d86d7bbaf Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Mon, 28 Sep 2026 16:00:16 +0200 Subject: [PATCH 17/18] docs: single backticks for inline code, as in the dycore --- .../tmx/scalar_diffusion.py | 18 +++++++++--------- .../tmx/stencils/scalar_diffusion.py | 4 ++-- .../subgrid_scale_physics/tmx/tmx_states.py | 8 ++++---- .../tmx/tests/tmx/fixtures.py | 2 +- .../test_tmx_scalar_diffusion.py | 2 +- 5 files changed, 17 insertions(+), 17 deletions(-) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py index 1ce33f9aeb..22c29fb105 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py @@ -8,8 +8,8 @@ """The scalar diffusion component of tmx. -Port of ``Compute_diffusion_hydrometeors`` and ``Compute_diffusion_temperature`` in ICON's -``mo_vdf.f90``, with the implicit vertical solver. +Port of `Compute_diffusion_hydrometeors` and `Compute_diffusion_temperature` in ICON's +`mo_vdf.f90`, with the implicit vertical solver. """ from __future__ import annotations @@ -57,7 +57,7 @@ def __init__( scale_turb_energy_flux: float, ) -> None: self._exchange = exchange - # ``zfactor`` in Compute_diffusion_temperature + # `zfactor` in Compute_diffusion_temperature energy_flux_factor = scale_turb_energy_flux if use_scale_turb_energy_flux else 1.0 use_internal_energy = energy_type == tmx_config.EnergyType.INTERNAL @@ -143,10 +143,10 @@ def run_hydrometeor_diffusion( dtime: float, ) -> None: """ - Diffuse qv, qc and qi (``Compute_diffusion_hydrometeors`` in mo_vdf.f90, without CO2). + Diffuse qv, qc and qi (`Compute_diffusion_hydrometeors` in mo_vdf.f90, without CO2). - Only qv has a surface flux, the evapotranspiration. Needs ``kh_ic`` and ``km_ie`` of - ``diagnostic_state``. + Only qv has a surface flux, the evapotranspiration. Needs `kh_ic` and `km_ie` of + `diagnostic_state`. """ log.debug("tmx hydrometeor diffusion (Compute_diffusion_hydrometeors): start") @@ -203,12 +203,12 @@ def run_temperature_diffusion( dtime: float, ) -> None: """ - Diffuse the temperature as dry static or internal energy (``Compute_diffusion_temperature`` + Diffuse the temperature as dry static or internal energy (`Compute_diffusion_temperature` in mo_vdf.f90). The energy is computed with the input tracers and converted back with the new qv, qc and - qi of ``new_state``, so this runs after ``run_hydrometeor_diffusion``. Needs - ``kh_ic`` and ``km_ie`` of ``diagnostic_state``. + qi of `new_state`, so this runs after `run_hydrometeor_diffusion`. Needs + `kh_ic` and `km_ie` of `diagnostic_state`. """ log.debug("tmx temperature diffusion (Compute_diffusion_temperature): start") diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py index c80379707f..f57417c472 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py @@ -218,7 +218,7 @@ def _compute_temperature_from_internal_energy( height_above_ground: fa.CellKField[wpfloat], grav: wpfloat, ) -> fa.CellKField[wpfloat]: - """Inverse of ``_compute_internal_energy_from_temperature``.""" + """Inverse of `_compute_internal_energy_from_temperature`.""" one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) return compute_temperature_from_internal_energy_per_area( energy - grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd, @@ -346,7 +346,7 @@ def _diffuse_energy_and_update_temperature( ) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: """ New temperature and its tendency after one diffusion step of the energy of - ``_compute_energy_from_temperature``, converted back with the new qv, qc and qi. + `_compute_energy_from_temperature`, converted back with the new qv, qc and qi. """ a, b, c = _assemble_scalar_diffusion_matrix( diffusivity, inv_dz, air_mass, prefactor, minlvl, maxlvl diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py index 73135fa704..c3e8a18424 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py @@ -87,9 +87,9 @@ class TmxSurfaceFluxState: """Surface fluxes provided by the surface scheme (inputs to the atmospheric diffusion).""" evapotranspiration: fa.CellField[ta.wpfloat] - """Surface evapotranspiration flux (``evspsbl``) [kg/(m^2 s)].""" + """Surface evapotranspiration flux (`evspsbl`) [kg/(m^2 s)].""" sensible_heat_flux: fa.CellField[ta.wpfloat] - """Surface sensible heat flux (``hfss``) [W/m^2].""" + """Surface sensible heat flux (`hfss`) [W/m^2].""" @dataclasses.dataclass(frozen=True) @@ -123,7 +123,7 @@ class TmxInputState: rho: fa.CellKField[ta.wpfloat] """Air density on full levels [kg/m^3].""" air_mass: fa.CellKField[ta.wpfloat] - """Air mass per unit area (``mair``) on full levels [kg/m^2].""" + """Air mass per unit area (`mair`) on full levels [kg/m^2].""" @dataclasses.dataclass(frozen=True) @@ -211,7 +211,7 @@ def allocate( @dataclasses.dataclass(frozen=True) class TmxNewState: - """Fields updated by the tmx diffusion: ``new = state + tend * dtime``.""" + """Fields updated by the tmx diffusion: `new = state + tend * dtime`.""" temperature: fa.CellKField[ta.wpfloat] """Updated air temperature on full levels [K].""" diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py index b6bb155fd0..483da8948c 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/fixtures.py @@ -47,7 +47,7 @@ def tmx_dtime( process_props: decomposition.ProcessProperties, download_ser_data: None, # downloads data as side-effect ) -> float: - """The tmx time step [s]: ``dt_vdf`` of the experiment's ``aes_phy_nml``.""" + """The tmx time step [s]: `dt_vdf` of the experiment's `aes_phy_nml`.""" input_dict = dt_utils.load_fortran_dict( experiment_description=experiment_description, process_props=process_props, diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py index fb7a222997..93359a7f04 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py @@ -8,7 +8,7 @@ """Integration tests of the tmx scalar diffusion component. -Both stages are seeded from the tmx-diagnostics-exit savepoint (``kh_ic``, ``km_ie``) and the +Both stages are seeded from the tmx-diagnostics-exit savepoint (`kh_ic`, `km_ie`) and the temperature stage also from the tmx-hydro-exit savepoint (the new qv, qc, qi), so that failures do not cascade between them. """ From 21364b1ff653d42961e99046e214d804fdae788a Mon Sep 17 00:00:00 2001 From: Hannes Vogt Date: Tue, 29 Sep 2026 18:07:22 +0200 Subject: [PATCH 18/18] Act on the third #1490 review round Remove the dry-static-energy path (EnergyType keeps only INTERNAL, with a clear error for 1), take the experiment dates from EXCLAIM_APE_AES.dates, add TODOs for the energy-function naming and the energy exchange, and drop the tmx_ prefix from the tmx integration and unit test file names. --- .../integration_tests/test_muphys_datatest.py | 5 +- .../subgrid_scale_physics/tmx/config.py | 13 +- .../tmx/scalar_diffusion.py | 12 +- .../tmx/stencils/scalar_diffusion.py | 91 +++--------- .../subgrid_scale_physics/tmx/tmx_states.py | 2 +- ...tmx_diagnostics.py => test_diagnostics.py} | 0 ...list_config.py => test_namelist_config.py} | 0 ..._diffusion.py => test_scalar_diffusion.py} | 1 - .../tmx/tests/tmx/integration_tests/utils.py | 10 +- .../stencil_tests/test_scalar_diffusion.py | 130 ++++++------------ .../{test_tmx_config.py => test_config.py} | 11 +- .../src/icon4py/model/testing/definitions.py | 2 + 12 files changed, 94 insertions(+), 183 deletions(-) rename model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/{test_tmx_diagnostics.py => test_diagnostics.py} (100%) rename model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/{test_tmx_namelist_config.py => test_namelist_config.py} (100%) rename model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/{test_tmx_scalar_diffusion.py => test_scalar_diffusion.py} (99%) rename model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/{test_tmx_config.py => test_config.py} (81%) diff --git a/model/atmosphere/subgrid_scale_physics/muphys/tests/muphys/integration_tests/test_muphys_datatest.py b/model/atmosphere/subgrid_scale_physics/muphys/tests/muphys/integration_tests/test_muphys_datatest.py index 2eac321518..a8afa41a43 100644 --- a/model/atmosphere/subgrid_scale_physics/muphys/tests/muphys/integration_tests/test_muphys_datatest.py +++ b/model/atmosphere/subgrid_scale_physics/muphys/tests/muphys/integration_tests/test_muphys_datatest.py @@ -62,10 +62,7 @@ "experiment_description", [definitions.Experiments.EXCLAIM_APE_AES], ) -@pytest.mark.parametrize( - "date", - ["2008-09-01T00:00:00.000", "2008-09-01T00:05:00.000", "2008-09-01T00:10:00.000"], -) +@pytest.mark.parametrize("date", definitions.Experiments.EXCLAIM_APE_AES.dates) def test_muphys_granule( date: str, *, diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py index 02cb30f8c5..e3439f6722 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/config.py @@ -30,9 +30,8 @@ class SolverType(int, enum.Enum): @config_io.register_enum class EnergyType(int, enum.Enum): - """Type of energy diffused by the temperature (heat) diffusion.""" + """Type of energy diffused by the heat diffusion; ICON's dry static energy (1) is not ported.""" - DRY_STATIC = 1 # dry static energy cp*T + g*z INTERNAL = 2 # internal energy cv*T @@ -53,7 +52,7 @@ class TmxConfig: energy_type: typing.Annotated[ EnergyType, common_conf_opt.ConfigOption( - description="Type of energy diffused by the heat diffusion (dry static or internal).", + description="Type of energy diffused by the heat diffusion (only internal energy is implemented).", ), ] = EnergyType.INTERNAL @@ -161,7 +160,13 @@ def __post_init__(self) -> None: f"Invalid argument 'solver_type': only the implicit solver " f"({SolverType.IMPLICIT.value}) is implemented, got {self.solver_type}." ) from None - self.energy_type = EnergyType(self.energy_type) + try: + self.energy_type = EnergyType(self.energy_type) + except ValueError: + raise ValueError( + f"Invalid argument 'energy_type': only internal energy " + f"({EnergyType.INTERNAL.value}) is implemented, got {self.energy_type}." + ) from None if self.turb_prandtl <= 0.0: raise ValueError( diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py index 22c29fb105..0b3d1e0826 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/scalar_diffusion.py @@ -20,7 +20,7 @@ import gt4py.next as gtx -from icon4py.model.atmosphere.subgrid_scale_physics.tmx import config as tmx_config, tmx_states +from icon4py.model.atmosphere.subgrid_scale_physics.tmx import tmx_states from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils import ( scalar_diffusion as scalar_stencils, ) @@ -52,14 +52,12 @@ def __init__( backend: model_backends.BackendLike, exchange: decomposition.ExchangeRuntime, turb_prandtl: float, - energy_type: tmx_config.EnergyType, use_scale_turb_energy_flux: bool, scale_turb_energy_flux: float, ) -> None: self._exchange = exchange # `zfactor` in Compute_diffusion_temperature energy_flux_factor = scale_turb_energy_flux if use_scale_turb_energy_flux else 1.0 - use_internal_energy = energy_type == tmx_config.EnergyType.INTERNAL cell_domain = h_grid.domain(dims.CellDim) cell_start_nudging = grid.start_index(cell_domain(h_grid.Zone.NUDGING)) @@ -110,7 +108,6 @@ def __init__( constant_args={ "height_above_ground": metric_state.height_above_ground, "grav": constants.GRAV, - "use_internal_energy": use_internal_energy, }, horizontal_sizes=horizontal_sizes, vertical_sizes=vertical_sizes, @@ -125,7 +122,6 @@ def __init__( "height_above_ground": metric_state.height_above_ground, "prefactor": energy_flux_factor, "grav": constants.GRAV, - "use_internal_energy": use_internal_energy, }, horizontal_sizes=horizontal_sizes, vertical_sizes=vertical_sizes, @@ -203,8 +199,8 @@ def run_temperature_diffusion( dtime: float, ) -> None: """ - Diffuse the temperature as dry static or internal energy (`Compute_diffusion_temperature` - in mo_vdf.f90). + Diffuse the temperature as internal energy (`Compute_diffusion_temperature` in + mo_vdf.f90). The energy is computed with the input tracers and converted back with the new qv, qc and qi of `new_state`, so this runs after `run_hydrometeor_diffusion`. Needs @@ -224,6 +220,8 @@ def run_temperature_diffusion( ) log.debug("communication of energy (cells): start") + # TODO(havogt): computing the energy inside the diffusion program, recomputed at the + # neighbour cells, would drop this exchange; try it once MPI tests check it. self._exchange.exchange(dims.CellDim, self.energy) log.debug("communication of energy (cells): end") diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py index f57417c472..9928342b7d 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/stencils/scalar_diffusion.py @@ -18,7 +18,6 @@ from icon4py.model.common.constants import PhysicsConstants from icon4py.model.common.dimension import C2E, E2C, C2EDim from icon4py.model.common.physics.thermodynamics.compute_energy import ( - _compute_dry_static_energy, compute_internal_energy_per_area, ) from icon4py.model.common.physics.thermodynamics.compute_temperature import ( @@ -192,19 +191,24 @@ def diffuse_tracer( ) +# TODO(havogt): the name omits the q tracers it depends on; revisit it, and its pair +# `_compute_temperature_from_internal_energy`, in a naming round. @gtx.field_operator def _compute_internal_energy_from_temperature( temperature: fa.CellKField[wpfloat], qv: fa.CellKField[wpfloat], - q_liquid: fa.CellKField[wpfloat], - q_solid: fa.CellKField[wpfloat], + qc: fa.CellKField[wpfloat], + qi: fa.CellKField[wpfloat], + qr: fa.CellKField[wpfloat], + qs: fa.CellKField[wpfloat], + qg: fa.CellKField[wpfloat], height_above_ground: fa.CellKField[wpfloat], grav: wpfloat, ) -> fa.CellKField[wpfloat]: """Specific internal energy plus cvd / cpd times the geopotential above ground.""" one = broadcast(wpfloat("1.0"), (dims.CellDim, dims.KDim)) return ( - compute_internal_energy_per_area(temperature, qv, q_liquid, q_solid, one, one) + compute_internal_energy_per_area(temperature, qv, qc + qr, qi + qs + qg, one, one) + grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd ) @@ -213,8 +217,11 @@ def _compute_internal_energy_from_temperature( def _compute_temperature_from_internal_energy( energy: fa.CellKField[wpfloat], qv: fa.CellKField[wpfloat], - q_liquid: fa.CellKField[wpfloat], - q_solid: fa.CellKField[wpfloat], + qc: fa.CellKField[wpfloat], + qi: fa.CellKField[wpfloat], + qr: fa.CellKField[wpfloat], + qs: fa.CellKField[wpfloat], + qg: fa.CellKField[wpfloat], height_above_ground: fa.CellKField[wpfloat], grav: wpfloat, ) -> fa.CellKField[wpfloat]: @@ -223,22 +230,13 @@ def _compute_temperature_from_internal_energy( return compute_temperature_from_internal_energy_per_area( energy - grav * height_above_ground * PhysicsConstants.cvd / PhysicsConstants.cpd, qv, - q_liquid, - q_solid, + qc + qr, + qi + qs + qg, one, one, ) -@gtx.field_operator -def _compute_temperature_from_dry_static_energy( - energy: fa.CellKField[wpfloat], - height_above_ground: fa.CellKField[wpfloat], - grav: wpfloat, -) -> fa.CellKField[wpfloat]: - return (energy - grav * height_above_ground) / PhysicsConstants.cpd - - @gtx.field_operator def _compute_surface_internal_energy_flux( sensible_heat_flux: fa.CellField[wpfloat], @@ -250,35 +248,6 @@ def _compute_surface_internal_energy_flux( ) -@gtx.field_operator -def _compute_surface_dry_static_energy_flux( - sensible_heat_flux: fa.CellField[wpfloat], -) -> fa.CellField[wpfloat]: - return sensible_heat_flux * PhysicsConstants.cpd / PhysicsConstants.cvd - - -@gtx.field_operator -def _compute_energy_from_temperature( - temperature: fa.CellKField[wpfloat], - qv: fa.CellKField[wpfloat], - qc: fa.CellKField[wpfloat], - qi: fa.CellKField[wpfloat], - qr: fa.CellKField[wpfloat], - qs: fa.CellKField[wpfloat], - qg: fa.CellKField[wpfloat], - height_above_ground: fa.CellKField[wpfloat], - grav: wpfloat, - use_internal_energy: bool, -) -> fa.CellKField[wpfloat]: - return ( - _compute_internal_energy_from_temperature( - temperature, qv, qc + qr, qi + qs + qg, height_above_ground, grav - ) - if use_internal_energy - else _compute_dry_static_energy(temperature, height_above_ground, grav) - ) - - @gtx.program(grid_type=gtx.GridType.UNSTRUCTURED) def compute_energy_from_temperature( temperature: fa.CellKField[wpfloat], @@ -291,13 +260,12 @@ def compute_energy_from_temperature( height_above_ground: fa.CellKField[wpfloat], energy: fa.CellKField[wpfloat], grav: wpfloat, - use_internal_energy: bool, horizontal_start: gtx.int32, horizontal_end: gtx.int32, vertical_start: gtx.int32, vertical_end: gtx.int32, ) -> None: - _compute_energy_from_temperature( + _compute_internal_energy_from_temperature( temperature, qv, qc, @@ -307,7 +275,6 @@ def compute_energy_from_temperature( qg, height_above_ground, grav, - use_internal_energy, out=energy, domain={ dims.CellDim: (horizontal_start, horizontal_end), @@ -342,29 +309,21 @@ def _diffuse_energy_and_update_temperature( dtime: wpfloat, minlvl: gtx.int32, maxlvl: gtx.int32, - use_internal_energy: bool, ) -> tuple[fa.CellKField[wpfloat], fa.CellKField[wpfloat]]: """ New temperature and its tendency after one diffusion step of the energy of - `_compute_energy_from_temperature`, converted back with the new qv, qc and qi. + `_compute_internal_energy_from_temperature`, converted back with the new qv, qc and qi. """ a, b, c = _assemble_scalar_diffusion_matrix( diffusivity, inv_dz, air_mass, prefactor, minlvl, maxlvl ) - # only the bottom row is used, where 'temperature' is that of the lowest level - surface_flux = ( - _compute_surface_internal_energy_flux(sensible_heat_flux, evapotranspiration, temperature) - if use_internal_energy - else broadcast( - _compute_surface_dry_static_energy_flux(sensible_heat_flux), (dims.CellDim, dims.KDim) - ) - ) new_energy, _ = _diffuse_scalar( energy, a, b, c, - surface_flux, + # only the bottom row is used, where 'temperature' is that of the lowest level + _compute_surface_internal_energy_flux(sensible_heat_flux, evapotranspiration, temperature), air_mass, rho, km_ie, @@ -375,14 +334,8 @@ def _diffuse_energy_and_update_temperature( dtime, maxlvl, ) - q_liquid = new_qc + qr - q_solid = new_qi + qs + qg - new_temperature = ( - _compute_temperature_from_internal_energy( - new_energy, new_qv, q_liquid, q_solid, height_above_ground, grav - ) - if use_internal_energy - else _compute_temperature_from_dry_static_energy(new_energy, height_above_ground, grav) + new_temperature = _compute_temperature_from_internal_energy( + new_energy, new_qv, new_qc, new_qi, qr, qs, qg, height_above_ground, grav ) return new_temperature, (new_temperature - temperature) * (wpfloat("1.0") / dtime) @@ -413,7 +366,6 @@ def diffuse_energy_and_update_temperature( prefactor: wpfloat, grav: wpfloat, dtime: wpfloat, - use_internal_energy: bool, horizontal_start: gtx.int32, horizontal_end: gtx.int32, vertical_start: gtx.int32, @@ -444,7 +396,6 @@ def diffuse_energy_and_update_temperature( dtime, vertical_start, vertical_end - 1, - use_internal_energy, out=(new_temperature, tend_temperature), domain={ dims.CellDim: (horizontal_start, horizontal_end), diff --git a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py index c3e8a18424..e5a7db686e 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/src/icon4py/model/atmosphere/subgrid_scale_physics/tmx/tmx_states.py @@ -211,7 +211,7 @@ def allocate( @dataclasses.dataclass(frozen=True) class TmxNewState: - """Fields updated by the tmx diffusion: `new = state + tend * dtime`.""" + """Fields updated by tmx: `new = state + tend * dtime`.""" temperature: fa.CellKField[ta.wpfloat] """Updated air temperature on full levels [K].""" diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_diagnostics.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_diagnostics.py similarity index 100% rename from model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_diagnostics.py rename to model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_diagnostics.py diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_namelist_config.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_namelist_config.py similarity index 100% rename from model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_namelist_config.py rename to model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_namelist_config.py diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_scalar_diffusion.py similarity index 99% rename from model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py rename to model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_scalar_diffusion.py index 93359a7f04..31d0274cc5 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_tmx_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/test_scalar_diffusion.py @@ -79,7 +79,6 @@ def _setup( backend=backend, exchange=decomposition.SingleNodeExchange(), turb_prandtl=config.turb_prandtl, - energy_type=config.energy_type, use_scale_turb_energy_flux=config.use_scale_turb_energy_flux, scale_turb_energy_flux=config.scale_turb_energy_flux, ) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py index 60fde7c82d..6145bbfb0d 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/integration_tests/utils.py @@ -18,6 +18,7 @@ from icon4py.model.atmosphere.subgrid_scale_physics.tmx import tmx_states from icon4py.model.common import dimension as dims from icon4py.model.common.metrics import metric_fields +from icon4py.model.testing import definitions if TYPE_CHECKING: @@ -26,11 +27,10 @@ from icon4py.model.testing import serialbox as sb -# Serialized timesteps of the exclaim_ape_aesPhys archive (run start -# 2008-09-01T00:00:00Z, dtime = 300 s). The archive also holds the -# 00:00:00 step, but that is the call made during model initialization, -# so the verification tests parametrize over the subsequent steps only. -TMX_DATES: tuple[str, ...] = ("2008-09-01T00:05:00.000", "2008-09-01T00:10:00.000") +# Serialized timesteps of the exclaim_ape_aesPhys archive. The first one is the +# call made during model initialization, so the verification tests parametrize +# over the subsequent steps only. +TMX_DATES: tuple[str, ...] = definitions.Experiments.EXCLAIM_APE_AES.dates[1:] # Relative tolerance of all tmx integration datatests. RTOL: float = 3.0e-12 diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py index a7f0594e34..8183aa4d0c 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/stencil_tests/test_scalar_diffusion.py @@ -11,7 +11,6 @@ import gt4py.next as gtx import numpy as np -import pytest from icon4py.model.atmosphere.subgrid_scale_physics.tmx.stencils.scalar_diffusion import ( compute_energy_from_temperature, @@ -53,16 +52,13 @@ def energy_from_temperature_numpy( q_solid: np.ndarray, height_above_ground: np.ndarray, grav: float, - use_internal_energy: bool, ) -> np.ndarray: - if use_internal_energy: - return ( - moist_heat_capacity_numpy(qv, q_liquid, q_solid) * temperature - - q_liquid * phy.lvc - - q_solid * phy.lsc - + grav * height_above_ground * phy.cvd / phy.cpd - ) - return phy.cpd * temperature + grav * height_above_ground + return ( + moist_heat_capacity_numpy(qv, q_liquid, q_solid) * temperature + - q_liquid * phy.lvc + - q_solid * phy.lsc + + grav * height_above_ground * phy.cvd / phy.cpd + ) def temperature_from_energy_numpy( @@ -73,14 +69,11 @@ def temperature_from_energy_numpy( q_solid: np.ndarray, height_above_ground: np.ndarray, grav: float, - use_internal_energy: bool, ) -> np.ndarray: - if use_internal_energy: - internal_energy = energy - grav * height_above_ground * phy.cvd / phy.cpd - return (internal_energy + q_liquid * phy.lvc + q_solid * phy.lsc) / ( - moist_heat_capacity_numpy(qv, q_liquid, q_solid) - ) - return (energy - grav * height_above_ground) / phy.cpd + internal_energy = energy - grav * height_above_ground * phy.cvd / phy.cpd + return (internal_energy + q_liquid * phy.lvc + q_solid * phy.lsc) / ( + moist_heat_capacity_numpy(qv, q_liquid, q_solid) + ) def diffuse_scalar_numpy( @@ -248,7 +241,6 @@ class TestComputeEnergyFromTemperature(stencil_tests.StencilTest): stencil_tests.StandardStaticVariants.COMPILE_TIME_DOMAIN: ( *_DOMAIN_ARGS, "grav", - "use_internal_energy", ), } @@ -265,7 +257,6 @@ def reference( qg: np.ndarray, height_above_ground: np.ndarray, grav: float, - use_internal_energy: bool, horizontal_start: int, horizontal_end: int, vertical_start: int, @@ -282,37 +273,23 @@ def reference( q_solid=q_solid, height_above_ground=height_above_ground, grav=grav, - use_internal_energy=use_internal_energy, )[cells, rows] return dict(energy=energy) - @stencil_tests.input_data_fixture( - params=[True, False], ids=["internal_energy", "dry_static_energy"] - ) - def input_data( - data_alloc: stencil_tests.DataAllocationWrapper, - grid: base.Grid, - request: pytest.FixtureRequest, - ) -> dict: - return _compute_energy_input_data(data_alloc, grid, use_internal_energy=request.param) - - -def _compute_energy_input_data( - data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, use_internal_energy: bool -) -> dict: - horizontal_start, horizontal_end = _cells(grid) - return dict( - temperature=data_alloc.random_field(dims.CellDim, dims.KDim, low=200.0, high=300.0), - **_tracers(data_alloc), - height_above_ground=data_alloc.random_field(dims.CellDim, dims.KDim, high=1.0e4), - energy=data_alloc.zero_field(dims.CellDim, dims.KDim), - grav=constants.GRAV, - use_internal_energy=use_internal_energy, - horizontal_start=horizontal_start, - horizontal_end=horizontal_end, - vertical_start=gtx.int32(0), - vertical_end=gtx.int32(grid.num_levels), - ) + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + horizontal_start, horizontal_end = _cells(grid) + return dict( + temperature=data_alloc.random_field(dims.CellDim, dims.KDim, low=200.0, high=300.0), + **_tracers(data_alloc), + height_above_ground=data_alloc.random_field(dims.CellDim, dims.KDim, high=1.0e4), + energy=data_alloc.zero_field(dims.CellDim, dims.KDim), + grav=constants.GRAV, + horizontal_start=horizontal_start, + horizontal_end=horizontal_end, + vertical_start=gtx.int32(0), + vertical_end=gtx.int32(grid.num_levels), + ) class TestDiffuseEnergyAndUpdateTemperature(stencil_tests.StencilTest): @@ -325,7 +302,6 @@ class TestDiffuseEnergyAndUpdateTemperature(stencil_tests.StencilTest): "rturb_prandtl", "prefactor", "grav", - "use_internal_energy", ), } @@ -350,7 +326,6 @@ def reference( prefactor: float, grav: float, dtime: float, - use_internal_energy: bool, horizontal_start: int, horizontal_end: int, vertical_start: int, @@ -364,13 +339,10 @@ def reference( 1.0 / air_mass[:, rows], ) a, b, c = matrix_diagonals_on_rows(matrix, air_mass.shape, rows) - if use_internal_energy: - temperature_sfc = temperature[:, vertical_end - 1] - surface_flux = sensible_heat_flux + temperature_sfc * evapotranspiration * ( - phy.cvv - phy.cvd - ) - else: - surface_flux = sensible_heat_flux * phy.cpd / phy.cvd + temperature_sfc = temperature[:, vertical_end - 1] + surface_flux = sensible_heat_flux + temperature_sfc * evapotranspiration * ( + phy.cvv - phy.cvd + ) new_energy, _ = diffuse_scalar_numpy( stencil_tests.connectivities_asnumpy(grid), var=energy, @@ -404,39 +376,25 @@ def reference( q_solid=q_solid, height_above_ground=height_above_ground, grav=grav, - use_internal_energy=use_internal_energy, )[cells, rows] tend_temperature[cells, rows] = ( new_temperature[cells, rows] - temperature[cells, rows] ) / dtime return dict(new_temperature=new_temperature, tend_temperature=tend_temperature) - @stencil_tests.input_data_fixture( - params=[True, False], ids=["internal_energy", "dry_static_energy"] - ) - def input_data( - data_alloc: stencil_tests.DataAllocationWrapper, - grid: base.Grid, - request: pytest.FixtureRequest, - ) -> dict: - return _diffuse_energy_input_data(data_alloc, grid, use_internal_energy=request.param) - - -def _diffuse_energy_input_data( - data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid, use_internal_energy: bool -) -> dict: - return dict( - **_diffusion_input_data(data_alloc, grid, "energy"), - diffusivity=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.0), - inv_dz=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.1), - **_tracers(data_alloc, prefix="new_"), - sensible_heat_flux=data_alloc.random_field(dims.CellDim), - evapotranspiration=data_alloc.random_field(dims.CellDim), - temperature=data_alloc.random_field(dims.CellDim, dims.KDim, low=200.0, high=300.0), - height_above_ground=data_alloc.random_field(dims.CellDim, dims.KDim, high=1.0e4), - new_temperature=data_alloc.zero_field(dims.CellDim, dims.KDim), - tend_temperature=data_alloc.zero_field(dims.CellDim, dims.KDim), - prefactor=wpfloat(1.5), - grav=constants.GRAV, - use_internal_energy=use_internal_energy, - ) + @stencil_tests.input_data_fixture + def input_data(data_alloc: stencil_tests.DataAllocationWrapper, grid: base.Grid) -> dict: + return dict( + **_diffusion_input_data(data_alloc, grid, "energy"), + diffusivity=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.0), + inv_dz=data_alloc.random_field(dims.CellDim, dims.KHalfDim, low=0.1), + **_tracers(data_alloc, prefix="new_"), + sensible_heat_flux=data_alloc.random_field(dims.CellDim), + evapotranspiration=data_alloc.random_field(dims.CellDim), + temperature=data_alloc.random_field(dims.CellDim, dims.KDim, low=200.0, high=300.0), + height_above_ground=data_alloc.random_field(dims.CellDim, dims.KDim, high=1.0e4), + new_temperature=data_alloc.zero_field(dims.CellDim, dims.KDim), + tend_temperature=data_alloc.zero_field(dims.CellDim, dims.KDim), + prefactor=wpfloat(1.5), + grav=constants.GRAV, + ) diff --git a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_tmx_config.py b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_config.py similarity index 81% rename from model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_tmx_config.py rename to model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_config.py index e41f308bb2..3bf6d32cd3 100644 --- a/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_tmx_config.py +++ b/model/atmosphere/subgrid_scale_physics/tmx/tests/tmx/unit_tests/test_config.py @@ -25,9 +25,9 @@ def test_config_rejects_negative_km_min() -> None: def test_config_coerces_enums_from_ints() -> None: - config = tmx_config.TmxConfig(solver_type=2, energy_type=1) + config = tmx_config.TmxConfig(solver_type=2, energy_type=2) assert config.solver_type is tmx_config.SolverType.IMPLICIT - assert config.energy_type is tmx_config.EnergyType.DRY_STATIC + assert config.energy_type is tmx_config.EnergyType.INTERNAL @pytest.mark.parametrize("solver_type", [1, 3]) @@ -36,9 +36,10 @@ def test_config_rejects_unimplemented_solver_types(solver_type: int) -> None: tmx_config.TmxConfig(solver_type=solver_type) -def test_config_rejects_invalid_enum_values() -> None: - with pytest.raises(ValueError): - tmx_config.TmxConfig(energy_type=3) +@pytest.mark.parametrize("energy_type", [1, 3]) +def test_config_rejects_unimplemented_energy_types(energy_type: int) -> None: + with pytest.raises(ValueError, match="energy_type"): + tmx_config.TmxConfig(energy_type=energy_type) def test_config_round_trips_through_config_io() -> None: diff --git a/model/testing/src/icon4py/model/testing/definitions.py b/model/testing/src/icon4py/model/testing/definitions.py index 3956e43005..63ecdd2367 100644 --- a/model/testing/src/icon4py/model/testing/definitions.py +++ b/model/testing/src/icon4py/model/testing/definitions.py @@ -177,6 +177,7 @@ class ExperimentDescription: long_name: str grid: GridDescription version: int + dates: tuple[str, ...] = () @dataclasses.dataclass @@ -218,6 +219,7 @@ class Experiments: long_name="EXCLAIM Aquaplanet experiment. JW IC and AES physics", grid=Grids.R02B04_GLOBAL, version=11, + dates=("2008-09-01T00:00:00.000", "2008-09-01T00:05:00.000", "2008-09-01T00:10:00.000"), ) MCH_CH_R04B09: Final = ExperimentDescription( name="exclaim_ch_r04b09_dsl",