Repository navigation
Preserve slow EMA decay in low-precision updates - #913
Open
sylvesterkaczmarek wants to merge 1 commit into
Open
sylvesterkaczmarek wants to merge 1 commit into
sylvesterkaczmarek wants to merge 1 commit into
Conversation
Signed-off-by: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.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.
Summary
Compute decay complements, floating-point EMA updates and debiasing in at least float32, then restore the existing stored-state and output dtypes. This prevents valid slow decays from rounding to one in float16/bfloat16. State keys, scan carry types, explicit wider initialization and existing integer-input behavior are retained.
Low-precision checkpoint state remains low precision, so ordinary accumulation rounding is not eliminated. This change is independent of the documentation-only edits in PR #888.
Reproduction
A decay of
0.9999rounds to one when converted to float16 or bfloat16.ExponentialMovingAveragethen returns NaNs for the first debiased update, or a zero non-debiased update. A constant input[1, 2]should instead have first debiased result[1, 2].Validation
python -m pytest -q haiku/_src/moving_averages_test.py haiku/_src/moving_averages_precision_test.py35 tests passed. Twenty-one new cases cover three dtypes and decay values, debiased/non-debiased scans against a quantized-state reference, unchanged preview state, warmup, parameter trees, explicit float32 initialization, mixed input/state dtypes and integer compatibility. The complete existing moving-average test module passes. Existing deprecated initialization warnings remain. Persistent-state precision and schema are not changed.
Negative control on unchanged main: 11 failed, 10 passed in 1.72s.
Tested on macOS CPU with real module imports. New-test formatting, scoped static checks, syntax checks and
git diff --checkpass. GPU/TPU and the full repository suite were not run. No dependency or workflow changes.