Skip to content

Center running variance batches along the requested reduction axes - #1221

Open
sylvesterkaczmarek wants to merge 3 commits into
NVIDIA:mainfrom
sylvesterkaczmarek:fix/running-variance-reduction-axis-broadcast
Open

sylvesterkaczmarek wants to merge 3 commits into
NVIDIA:mainfrom
sylvesterkaczmarek:fix/running-variance-reduction-axis-broadcast

Conversation

@sylvesterkaczmarek

Copy link
Copy Markdown

Earth2Studio Pull Request

Description

Fixes #1220.

Keep reduced dimensions as singleton axes when centering each running-variance batch, then squeeze only those axes before storing its sum. The dropped-axis mean previously broadcast against the wrong dimensions: reducing time in a 3-by-3 tensor returned variances [17.5, 4, 17.5] instead of [1, 1, 1], while a 2-by-3 tensor raised a shape error. Running standard deviation inherited the same problem.

Preserve the reduced shape of stored sums, output coordinates, weighted normalization, singleton behavior and batch-combination formula. Update the changelog. No public API changes.

Validation

Sixteen parametrized cases extend the existing moments tests. They cover variance/std, float32/float64, square and nonsquare non-leading reductions, multiple-axis layouts, unequal streaming batch sizes, unchanged inputs, coordinate order, weighted centering and single-batch autograd/finite-difference checks. Streamed unweighted results are checked against direct Torch variance/std over all samples seen so far. The weighted control deliberately retains the existing running-weight denominator.

  • New tests against unchanged source: 16 failed.
  • Existing CPU moments suite: 20 passed.
  • python -m pytest -o addopts='' test/statistics/test_moments.py -k 'not cuda' -q: 36 passed, with 17 CUDA cases deselected.
  • All applicable pre-commit hooks passed on the changed files, including Ruff, mypy, Interrogate and Markdown checks. Compilation and git diff --check passed.

Testing used Python 3.12.11, PyTorch 2.14.1, on macOS CPU. The complete statistics/repository suites, GPU execution, external datasets and forecasting models were not tested. Existing optional-dependency, configuration and indexing warnings remain. Default batch-count ratios retain their existing float32 arithmetic, so streaming comparisons allow that rounding; the precision policy and weighted-estimator normalization are outside this fix. No dependencies, workflows or existing test expectations change.

Checklist

  • I am familiar with the Contributing Guidelines.
  • New and existing tests cover these changes.
  • The documentation is consistent with the unchanged public interface.
  • CHANGELOG.md is updated.
  • An issue is linked to this pull request.
  • Assess and address Greptile feedback.

Dependencies

None.

Based on main at e916317a7e53ee73cc8cf5cd97064e970cfd723e.

Signed-off-by: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.com>
@copy-pr-bot

copy-pr-bot Bot commented Oct 3, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@greptile-apps

greptile-apps Bot commented Oct 3, 2026 •

Copy link
Copy Markdown
Contributor

Disclaimer: This is AI-generated, please review response for accuracy

RetriggerConfidence Score: 4/5

[Medium risk] Fixes running variance calculation for non-leading reduction axes.

The PR appears safe to merge; extending the weighted streaming test would strengthen regression coverage.

Findings

  1. P2 Weighted streaming coverage is missing ▶

Summary

The PR keeps reduction axes during running-variance batch centering, then removes them before storing the sum. It adds variance and standard-deviation coverage for non-leading and multiple-axis reductions and updates the changelog.

Reviews (1) · Last reviewed commit: "Center running variance batches along th..."

Comment thread test/statistics/test_moments.py Outdated
@sylvesterkaczmarek
sylvesterkaczmarek force-pushed the fix/running-variance-reduction-axis-broadcast branch from 6989031 to ca230aa Compare October 4, 2026 00:25

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.

Running variance centers batches on the wrong axes

1 participant