Skip to content

Stabilize zero-inflated variance at large base means - #3487

Open
AHMETHAKANBEZIR1 wants to merge 1 commit into
pyro-ppl:devfrom
AHMETHAKANBEZIR1:fix/zero-inflated-variance
Open

AHMETHAKANBEZIR1 wants to merge 1 commit into
pyro-ppl:devfrom
AHMETHAKANBEZIR1:fix/zero-inflated-variance

Conversation

@AHMETHAKANBEZIR1

Copy link
Copy Markdown

Proposed changes

Fixes #3486. Replace subtraction of squared moments in ZeroInflatedDistribution.variance with the equivalent law-of-total-variance expression:

(1 - gate) * base_variance + (gate * base_mean) * mean

This avoids cancellation when the base mean dominates its variance. Multiplying gate * base_mean before the remaining mean also avoids an unnecessary overflowing mean square, e.g. a float32 Poisson rate of 1e20 with gate=1e-20 has variance approximately 2e20. The existing expression produces NaN. A rate of 1e10 with gate=0 currently produces zero rather than the Poisson variance of 1e10. This shared property applies to generic zero-inflated distributions as well as ZIP/ZINB.

New and existing tests

One focused parametrized regression uses Poisson and Normal bases in float32/float64, including ordinary, zero-gate, small-gate and unit-gate cases in a batch. Expected values use an independent double-precision conditional-variance calculation; gradients are compared with the analytic derivative with respect to the base mean parameter.

  • Exact unmodified dev ae65fa4 variance restored in memory: all 4 added cases fail (247 deselected). The final regression is used; no repository source is modified by this baseline harness.
  • Final full tests/distributions/test_zero_inflated.py: 251 passed on Python 3.12/PyTorch 2.10 CPU and 251 passed on Python 3.14/PyTorch 2.13 CPU.
  • Full repository Ruff check, Ruff format check across pyro/examples/tests/scripts/profiler, copyright-header check, and mypy across pyro/scripts/tests pass (553 source files). Header checking requires Python UTF-8 mode on this Windows host; the initial host cp1254 decoding failed before completing the scan, and the UTF-8 rerun passed. No source headers were changed.
  • Focused Sphinx autodoc build of the three public zero-inflated classes passes with warnings treated as errors. Existing docstrings are unchanged.
  • Not run: full package tests, full project docs/tutorials, GPU, compilation, or real inference/training benchmarks. This does not claim to make truly infinite base moments finite.

AI assistance disclosure

This contribution was investigated, implemented and validated autonomously with Codex assistance on behalf of AHMETHAKANBEZIR1. It has not received independent human code review. Codex is recorded as a co-author. Maintainer review is requested; local results are not a claim of upstream CI success.

Use the law of total variance without subtracting squared means.
Multiply by the gate before the remaining mean factor to avoid
overflow of an unnecessary intermediate square.

Co-authored-by: Codex <noreply@openai.com>

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.

ZeroInflatedDistribution variance loses finite values at large means

1 participant