Skip to content

Compute standardization stats in float64 - #759

Open
RudraDudhat2509 wants to merge 2 commits into
mllam:mainfrom
RudraDudhat2509:fix/standardization-stats-float64
Open

RudraDudhat2509 wants to merge 2 commits into
mllam:mainfrom
RudraDudhat2509:fix/standardization-stats-float64

Conversation

@RudraDudhat2509

@RudraDudhat2509 RudraDudhat2509 commented Sep 25, 2026 •

Copy link
Copy Markdown
Contributor

Describe your changes

The standardization std was sqrt(E[x²] - E[x]²) from float32 moments, which cancels when the mean is large compared to the std (mean 1e5, std 100 gave 90.5, std 5 gave NaN, a constant field gave 45 instead of 0). Details, the simulation and how NumPy, scikit-learn, PyTorch, anemoi-datasets and mllam-data-prep handle it are in #758.

  • sample_moments: per-sample mean and mean of squares in float64, one sample at a time so memory stays at one sample. Used for the parameter and diff passes, flux stats use float64 too.
  • _std_from_moments: clamps rounding-level negative variance to 0 and raises on clearly negative variance, instead of a NaN.
  • Stats are still saved as float32, since the diff pass applies them to float32 batches.

Measured on a MEPS-sized batch (4 x 65 x 63784 x 17, CPU): 1.05 s to 1.20 s per batch, lower peak memory for large batches.

10 new tests, 5 of them fail on the old computation (std 100, 20 and 5, a constant field, and inconsistent moments). main() itself isn't covered since it needs MEPS data. The two lines in main overlap with #748.

No new dependencies required.

Issue Link

closes #758

Type of change

  • 🐛 Bug fix (non-breaking change that fixes an issue)
  • ✨ New feature (non-breaking change that adds functionality)
  • 💥 Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • 📖 Documentation (Addition or improvements to documentation)

Checklist before requesting a review

  • My branch is up-to-date with the target branch - if not update your fork with the changes from the target branch (use pull with --rebase option if possible).
  • I have performed a self-review of my code
  • For any new/modified functions/classes I have added docstrings that clearly describe its purpose, expected inputs and returned values
  • I have placed in-line comments to clarify the intent of any hard-to-understand passages of my code
  • I have updated the README to cover introduced code changes
  • I have added tests that prove my fix is effective or that my feature works
  • I have given the PR a name that clearly describes the change, written in imperative form (context).
  • I have requested a reviewer and an assignee (assignee is responsible for merging). This applies only if you have write access to the repo, otherwise feel free to tag a maintainer to add a reviewer and assignee.

Author checklist after completed review

  • I have added a line to the CHANGELOG describing this change, in a section
    reflecting type of change (add section where missing):
    • added: when you have added new functionality
    • changed: when default behaviour of the code has been changed
    • fixes: when your contribution fixes a bug
    • maintenance: when your contribution is relates to repo maintenance, e.g. CI/CD or documentation

RudraDudhat2509 and others added 2 commits September 26, 2026 00:08
The std was sqrt(E[x^2] - E[x]^2) from float32 moments, which cancels when
the mean is large compared to the std: a variable with mean 1e5 and std 100
got std 90.5, std 20 got 32.0, std 5 got NaN, and a constant field got 45
instead of 0. The precision is already lost when main() averages x and x**2
in float32, so it cannot be fixed in save_stats alone.

Accumulate the per-sample moments in float64 (one sample at a time to keep
memory at one sample), clamp rounding-level negative variance to zero and
raise on clearly negative variance, and still save float32 so the later
(x - mean) / std on float32 batches is not upcast.

Co-authored-by: Claude Sonnet 5 <noreply@anthropic.com>
Co-authored-by: Claude Sonnet 5 <noreply@anthropic.com>
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.

compute_standardization_stats writes a wrong or NaN std when a variable's mean is large compared to its std

1 participant