Repository navigation
Compute standardization stats in float64 - #759
Open
RudraDudhat2509 wants to merge 2 commits into
Open
RudraDudhat2509 wants to merge 2 commits into
RudraDudhat2509 wants to merge 2 commits into
Conversation
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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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.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 inmainoverlap with #748.No new dependencies required.
Issue Link
closes #758
Type of change
Checklist before requesting a review
pullwith--rebaseoption if possible).Author checklist after completed review
reflecting type of change (add section where missing):