Skip to content

Accumulate low-precision search averages in float32 - #128

Open
sylvesterkaczmarek wants to merge 1 commit into
google-deepmind:mainfrom
sylvesterkaczmarek:fix/low-precision-search-backup-20261007
Open

sylvesterkaczmarek wants to merge 1 commit into
google-deepmind:mainfrom
sylvesterkaczmarek:fix/low-precision-search-backup-20261007

Conversation

@sylvesterkaczmarek

Copy link
Copy Markdown

Summary

Fixes #127.

Accumulate the node's weighted return sum and visit denominator in at least float32, then cast back to the existing tree dtype. This avoids float16 overflow before division and conversion of large visit counts to half precision.

Tree storage, visit accounting and float32 arithmetic are preserved. This is independent of #126's Q-normalization change.

Validation

python -m pytest -q --pyargs mctx

46 tests passed, including all 24 existing tests and 22 new cases. The new suite covers float16/bfloat16/float32 values, opposite-sign returns, large visit counts, eager/JIT execution and the public MuZero policy. Ten cases fail on unchanged source.

The full suite ran from a temporary child directory so the repository's existing JSON fixture paths resolve correctly. Another 22 tests passed in a separate compatibility checkout containing #116's backward-pass refactor with the arithmetic correction applied inside it; that refactor is not included here.

Tests used macOS CPU and JAX 0.11.2. No accelerator, trained model or full game benchmark was run. Already-overflowed returns and general float32 overflow remain outside this correction.

New-test formatting, scoped syntax/lint checks and git diff --check pass. No dependency or workflow changes.

Signed-off-by: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.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.

Finite float16 returns overflow during search-node averaging

1 participant