Skip to content

Add single precision option to the dycore (intermediate state) - #970

Open
starkphi wants to merge 169 commits into
mainfrom
add_single_precision_dycore_part1
Open

starkphi wants to merge 169 commits into
mainfrom
add_single_precision_dycore_part1

Conversation

@starkphi

@starkphi starkphi commented Dec 5, 2025 •

Copy link
Copy Markdown

Adds the option to set wp/vpfloat to single via the env var ICON4PY_FLOAT_PRECISION (previous options were 'double' and 'mixed') to run in single precision.

This PR was created to later reduce merge complexity of #886.

Changes

  • Replace native python float types by either wpfloat or vpfloat
    • Change type annotations where there were native Python floats were present (which are by default passed as gtx.float64 to gtx.field_operators and gtx.programs) and add casts around float literals
    • Swap wpfloat <-> vpfloat where an inconsistency with an operator was caught (to hopefully make fixing mixed precision easier in the future)
    • Replace DBL_EPS by WP_EPS and VP_EPS
  • Setup pytests for single precision
    • Add pytest marker single_precision_ready (running a test with ICON4PY_FLOAT_PRECISION=single deselects all tests without) and mark integration tests for the standalone driver and most components (required much larger tolerances for some cases)
    • Remove pytest option --enable-mixed-precision (wpfloat or vpfloat are already set according to ICON4PY_FLOAT_PRECISION at import time of a pytest file)
    • Added PRECISION_VARIANTS as dimension to the ci pipeline matrix
  • Keep calculations and results in gtx.float64 internally in geometry, metrics or interpolation factories independently of the selected precision. Mainly to keep the precision of the RBF and other geometry factors acceptable when running in ICON4PY_FLOAT_PRECISION
    • factory.store_allfloats_as_double forces float allocations in a FieldSource to double-precision
    • FieldSource.export_field casts to the dtype declared in the field metadata where the fields leave the factories

@starkphi
starkphi force-pushed the add_single_precision_dycore_part1 branch from dd41784 to adcf8fb Compare December 5, 2025 15:49
@starkphi
starkphi force-pushed the add_single_precision_dycore_part1 branch from 577d435 to f0d747e Compare December 8, 2025 11:08
@starkphi
starkphi force-pushed the add_single_precision_dycore_part1 branch from f0d747e to 5daea04 Compare December 8, 2025 14:30
@muellch

muellch commented Dec 8, 2025

Copy link
Copy Markdown
Contributor

cscs-ci run default

Comment thread model/driver/src/icon4py/model/driver/icon4py_driver.py Outdated

@egparedes egparedes left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It looks good in general although I have a few comments.

Comment thread model/common/src/icon4py/model/common/grid/vertical.py Outdated
Comment thread model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/solve_nonhydro.py Outdated
Comment thread model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/solve_nonhydro.py Outdated
Comment thread model/testing/src/icon4py/model/testing/test_utils.py Outdated
Comment thread model/atmosphere/dycore/tests/dycore/integration_tests/test_velocity_advection.py Outdated
Comment thread model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/velocity_advection.py Outdated
Comment thread model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/velocity_advection.py Outdated
Comment thread model/common/src/icon4py/model/common/type_alias.py Outdated
Comment thread model/testing/src/icon4py/model/testing/pytest_hooks.py
Comment thread model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/solve_nonhydro.py Outdated
Comment thread model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/solve_nonhydro.py Outdated
Comment thread model/atmosphere/dycore/src/icon4py/model/atmosphere/dycore/solve_nonhydro.py Outdated
Comment thread model/atmosphere/diffusion/src/icon4py/model/atmosphere/diffusion/diffusion.py Outdated
@starkphi

Copy link
Copy Markdown
Author

cscs-ci run default

@starkphi
starkphi force-pushed the add_single_precision_dycore_part1 branch from 8bf1c5b to f169bfe Compare September 22, 2026 05:11
@starkphi

Copy link
Copy Markdown
Author

cscs-ci run default;BACKENDS=gtfn_cpu;SESSIONS=model;MODEL_SUBSETS=datatest

@starkphi

Copy link
Copy Markdown
Author

cscs-ci run default;BACKENDS=gtfn_cpu;SESSIONS=model;MODEL_SUBSETS=datatest

@starkphi

Copy link
Copy Markdown
Author

cscs-ci run default;BACKENDS=gtfn_cpu;SESSIONS=model;MODEL_SUBSETS=datatest;FLOAT_PRECISIONS=single

@starkphi

Copy link
Copy Markdown
Author

cscs-ci run default;BACKENDS=gtfn_cpu;SESSIONS=model;MODEL_SUBSETS=datatest;FLOAT_PRECISIONS=single:double

@muellch muellch left a comment •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude found some bugs and issues.
I am not sure you want to fix all of them before merging, but at least the bugs I assume.

SNOW_INTERCEPT_PARAMETER_MMB8 = 0.000594
SNOW_INTERCEPT_PARAMETER_MMB9 = 0.000000
SNOW_INTERCEPT_PARAMETER_MMB10 = -0.003577
SNOW_INTERCEPT_PARAMETER_MMB1 = ta.wpfloat(0.476221)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I guess I am not happy that all constants now have to be declared ta.wpfloat(). Longer term a better solution is needed I think.

@starkphi starkphi Sep 29, 2026 •

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed, it's noisy. It's needed because a Python float in a field operator is typed as float64, and GT4Py doesn't silently narrow that to float32 fields. Longer term, I see two ways: (a) GT4Py treating Python scalar literals as "weak" (taking on the field's dtype, like NumPy's NEP 50), which would remove the wrappers entirely, or (b) keeping the constants in float64 (as the host/factory code needs anyway, cf. the vertical grid comment) and casting them once to wp at the boundary instead of per constant. I'd rather not do this in this PR though.

(answer written by Claude 🤖)

actual[:, jks:], reference[:, jks:], atol=1e-13, err_msg=f"{name} in cloud"
actual[:, jks:],
reference[:, jks:],
atol=1e-13 if test_utils.wp_is_dp else 3e-9,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Uglier to read, but nice if the datatests also pass for single with smaller tolerances now.

Comment on lines 235 to 236
eta_v_at_edge = gtx.astype(eta_v_at_edge_dp, ta.wpfloat)
exchange.exchange(dims.EdgeDim, eta_v_at_edge)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: gtx.astype copies, so this exchanges a throwaway copy while vn is computed from the un-exchanged eta_v_at_edge_dp (L248). I think vn halo edges are now wrong in multi-rank runs. Exchange eta_v_at_edge_dp directly and drop the copy?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Exchange now on eta_v_at_edge_dp directly, copy dropped (changed in f6bb879).

this_metadata = self.metadata[field_name]
if "dims" not in this_metadata or not this_metadata["dims"]:
raise TypeError(
f"This function is intended to return a Field. Field name {field_name!r} looks like a Scalar ('dims' missing in metadata)."

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: Two callers still use get() for scalars and now hit this TypeError: test_parallel_metrics.py:250 (NFLAT_GRADP) and test_benchmark_solve_nonhydro.py:81 (MEAN_CELL_AREA, in the bencher CI selection). Use get_scalar() for the first and get_wpfloat() for the second, as driver_utils.py:236 does.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Both callers use get_scalar(); get_wpfloat/get_int32 removed entirely (changed in 54fa65c).

)
dtype_metadata = this_metadata.get("dtype", ta.wpfloat)
# `astype` is a `BuiltInFunction`, whose overloads are erased by the decorator.
return cast("state_utils.GTXFieldType", gtx.astype(field, dtype_metadata))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: gtx.astype copies even when the dtype already matches, so every get() returns a fresh copy, in double too. That means extra memory for every static field and no shared cached buffer. Please only cast when needed, e.g. field if field.dtype.scalar_type == dtype_metadata else gtx.astype(field, dtype_metadata).

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added data_alloc.astype_if_needed, used in get() (changed in e787b2d).

#: ice crystal number concentration at threshold temperature for mixed-phase cloud
NIMIX = 5.0 * math.exp(
0.304 * (PhysicsConstants.tmelt - THRESHOLD_FREEZE_TEMPERATURE_MIXEDPHASE)
NIMIX = ta.wpfloat(5.0) * gtx.exp(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: On scalars, gtx.exp/gtx.gamma go through numpy and scipy.special.gamma (if importable), unlike math.* on main. This shifts NIMIX and several graupel coefficients by 1–2 ulp in double, and makes them depend on whether scipy is installed. Keep math.* and cast to wpfloat?

@starkphi starkphi Sep 29, 2026 •

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Intended: using gtx math builtins so that the constants are computed with the same implementations as when these functions are applied to fields. The resulting 1–2 ulp differences in double (NIMIX, CCSDEP) are below the differences to our actual reference (Fortran libm and compiler reordering), so we don't consider them relevant. scipy is in our lock file (and it's what gt4py[standard] uses), so the scipy-dependent path is the one that's always taken in practice.

(answer written by Claude 🤖)

qnc: float = 100.0

def __post_init__(self):
ta.dataclass_float_to_wp(self)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: In single, dt becomes np.float32, which datetime.timedelta(seconds=...) in test_component_datatest.py:69 rejects. Not hit yet, since that test isn't single-marked. timedelta(seconds=float(experiment.dt)) there fixes it.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Using timedelta(seconds=float(experiment.dt)) now (changed in 9706af2).

self,
"_vct_a",
vct_a,
gtx.astype(vct_a, ta.wpfloat),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: This breaks the "factories compute in float64" contract. vct_a/vct_b are stored as wpfloat, VerticalGridConfig is cast to wp, and all constants are wpfloat. So in single the metrics come from f32-rounded inputs (metrics_factory.py:196 only widens), and nflatlev/kstart_moist are decided on f32 heights. Suggestion: keep configs, vct_a/vct_b and host-side constants in f64, and cast to wp only in get(). To pin the contract, run the existing factory datatests in single against get_full_precision() with the double tolerances.

@starkphi starkphi Sep 29, 2026 •

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed, fixed: VerticalGridConfig is no longer cast to wp (all consumers are f64 factory code or host code), vct_a/vct_b are computed, read (file and savepoint) and stored in float64, and the level indices are determined on the f64 heights. interface_physical_height provides a cached wp version for the granules; the metrics factory uses vertical_grid.vct_a directly. VerticalGrid.__post_init__ now asserts float64 instead of casting (this also removed the InitVar/__setattr__ indirection; separate commit 4da177f). Measured before the fix: in single, vct_a computed in f32 had layer-thickness errors up to ~3e-5 relative, i.e. 30-500x the final cast to f32. (@egparedes: this also replaces the astype you asked about in __post_init__.) Constants: see my reply to the ta.wpfloat() comment.

(answer written by Claude 🤖)

attributes = [
field.name
for field in self.__dataclass_fields__.values()
if "float" in repr(field.type)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: Fragile: with from __future__ import annotations (all configs use it), field.type is the whole Annotated[...] source string, description included. A bool/int field whose description mentions "float" gets converted. None silently becomes NaN, and later assignments skip the cast. Use typing.get_type_hints() with a real type check, or an explicit field list?

@starkphi starkphi Sep 29, 2026 •

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

dataclass_float_to_wp now resolves the field types with typing.get_type_hints instead of matching on the annotation string, so descriptions can't trigger it anymore, and only dataclass fields whose resolved type is a float type are cast. Unions (e.g. wpfloat | None) are deliberately not cast: pass attributes explicitly or convert at the point of use (like CellParams.mean_cell_area). The tests that assigned config values after construction now pass them to the constructor. DiffusionConfig and NonHydrostaticConfig are now frozen as well (like the other configs that use the helper), so values can't be assigned after construction anymore.

(answer written by Claude 🤖)

)

assert test_utils.dallclose(
test_utils.assert_dallclose( # this is completely off in single! At least some are by factors of 100 larger

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤖 Claude: Worrying, given the test is marked single_precision_ready. Several single tolerances here are loose: rtol 5e-2 on the horizontal pressure gradient, 3e-2 on nonhydro_buoy, 1e-1 on vol_flx_ic, atol 3 on z_theta_v_fl_e. Should these tests guarantee no regression, or that single is correct? For the latter, a reference from an ICON wp=sp build would be more convincing.

@starkphi starkphi Sep 29, 2026 •

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

These datatests compare single-precision icon4py against the double-precision ICON reference, so in single they check that we stay within the expected rounding distance of the double reference (regression / sanity), not that single is "correct" by itself. We replaced the large single rtols by atols set from the measured maximum absolute errors: the violations were always at values near zero in cancellation-heavy quantities (vertical pressure gradient terms, z_th_ddz_exner_c, vertical fluxes), where a relative tolerance is meaningless. Relative to the largest value of each field, the atols are small: <= 2e-5 for the prognostic variables and fluxes (e.g. atol 10 for z_theta_v_fl_e, whose values reach 3.4e6, i.e. 3e-6 relative), set to about 2x the measured maximum error, and at most a few percent only for fields that are essentially zero (marked with a comment where the field is practically zero in APE). The stale comment is removed. The stronger statement is test_driver, which is single-ready and compares the final prognostic state after a full run (e.g. JW: vn within 1.5e-4 m/s, exner/theta_v/rho within ~1e-6 relative). A wp=sp ICON reference would indeed be more convincing. ICON can be built in single, but we're leaving serialized single-precision reference data out for now, because it would add a lot of additional data to store and load. Longer term these tolerances will probably be replaced by automatically tuned per-variable tolerances (cf. #1357).

(answer written by Claude 🤖)

Philipp Stark and others added 17 commits September 24, 2026 15:37
gtx.astype makes a copy with the current gt4py version. numpy/cupy would provide a non-copying version of the used command (xp.astype) but because the corresponding copy=True kwarg is not exposed to gt4py we work around it with a helper
Matches ICON's -HUGE(0._vp), so the fill value never wins the MAX in
enhance_diffusion_coefficient_for_grid_point_cold_pools. With VP_EPS,
kh_smag_e was clamped to >= 1.2e-7 in single precision.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
In single precision, MuphysExperiment.dt is cast to np.float32, which
datetime.timedelta rejects.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Missing annotations defaulted to gtx.float64, so names not in the
function signature passed validation. Replace the untyped sqrt lambda in
GridGeometry by a typed local function instead.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Resolves the actual type now instead of checking the repr string.

__post_init__ casting only works when the config is already constructed with all relevant variables. Therefore the config changes in test_diffusion were inserted into the constructor and the config dataclasses were changed to frozen=True.

(change inspired by [muellch's comment](#970 (comment)))
VerticalGridConfig is no longer cast to wpfloat, vct_a/vct_b are
computed, read and stored in float64, and the nflatlev/kstart_moist/
damping indices are determined on float64 heights. This restores the
"factories compute in float64" contract: in single precision, vct_a used
to be computed in float32 (layer thickness errors up to ~3e-5 relative).
interface_physical_height provides vct_a in working precision for the
granules; the metrics factory uses vertical_grid.vct_a directly.
Savepoint vct_a/vct_b are read in float64.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
All callers already pass float64, so the InitVar/private attribute/
__setattr__ indirection is not needed anymore. __post_init__ asserts
float64 instead of casting silently. test_damping_layer_calculation
built an int64 vct_a, which the cast used to hide.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
The existing factory tests are all datatests and not single precision
ready, so the precision-dependent part of the factories (casting the
float64 fields to the metadata dtype on export) was not covered in
single precision.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
The savepoint vct_a is read in float64 since the vertical grid keeps it
in double precision; the program expects working precision.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
@github-actions

Copy link
Copy Markdown

When developing, you can test your changes on CSCS CI before merge with the default pipeline: cscs-ci run default. This will run a default subset of tests.

You can pass options to override pipeline variables, for example:

  • cscs-ci run default;BACKENDS=gtfn_cpu;LEVELS=unit
  • cscs-ci run default;MODEL_SUBPACKAGES=common:driver;SESSIONS=model

Avoid running the pipeline for all tests when you are developing.

Available options are:

  • SESSIONS: model, model_mpi, or tools (correspond to nox sessions)
  • MODEL_SUBSETS: datatest, basic, or stencils (correspond to nox session selections)
  • MODEL_SUBPACKAGES: subpackages for non-MPI tests (last component, e.g. diffusion, driver)
  • MODEL_MPI_SUBPACKAGES: subpackages for MPI tests (as above)
  • BACKENDS: backends
  • GRIDS: grids for stencil tests (simple, icon_regional, or icon_global)
  • LEVELS: testing level for non-stencil tests (unit or integration)

For each option, all can be used as a shorthand for all possible values of that variable, e.g. LEVELS=all.

Multiple values can be given to each option with : used as the separator (; separates options and , separates pipelines).

See scripts/python/generate_ci_pipeline.py and noxfile.py for available values for each option.

The all pipeline can be run with cscs-ci run all. This will run all icon4py tests in CSCS CI which can be expensive. This pipeline runs on a schedule on main, and can be run when extensive validation is needed (e.g. before releases).

Merging

Once your PR is approved and ready for merging, add it to the merge queue. The merge CSCS CI pipeline will run automatically on the merge-queue branch and must pass before the PR is merged. A dummy merge check will be triggered on the PR itself since it's required to add a PR to the merge queue.

Optional Tests

To run benchmarks you can use:

  • cscs-ci run benchmark-bencher

For more detailed information please look at CI in the EXCLAIM universe.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants