Skip to content

Accumulate low-precision average pooling in float32 - #914

Open
sylvesterkaczmarek wants to merge 1 commit into
google-deepmind:mainfrom
sylvesterkaczmarek:fix/average-pool-low-precision-accumulation-20261006
Open

sylvesterkaczmarek wants to merge 1 commit into
google-deepmind:mainfrom
sylvesterkaczmarek:fix/average-pool-low-precision-accumulation-20261006

Conversation

@sylvesterkaczmarek

Copy link
Copy Markdown

Summary

Accumulate float16/bfloat16 window sums and SAME-padding counts in float32, then return the original dtype after division. This prevents premature half-precision sum overflow when the pooled mean is representable and improves bfloat16 averaging.

Float32/float64 inputs, max pooling, shape inference and padding semantics are unchanged. The change is separate from the shape-handling proposal in #824. Float32 accumulator overflow and accelerator performance are outside this correction.

Reproduction

import jax.numpy as jnp
import haiku as hk
x = jnp.full((1, 16, 16, 1), 256., dtype=jnp.float16)
print(hk.avg_pool(x, (1, 16, 16, 1), 1, 'VALID'))

Unchanged main returns infinity instead of 256.

Validation

python -m pytest -q haiku/_src/pool_test.py haiku/_src/pool_precision_test.py

All 20 existing pooling tests and 16 new cases pass. Coverage includes independent NumPy window means, positive/negative constant activations, nonuniform values, both layouts and padding modes, float16/bfloat16/float32, eager/JIT execution, vmap, and a real AvgPool module with a finite loss and gradients. Eight new tests fail on unchanged upstream and eight controls pass.

Eight existing shape-inference deprecation warnings were emitted. No trained model, image dataset or GPU/TPU execution was used; no performance claim is made.

Tested locally on macOS CPU using real module imports. New test formatting, scoped static checks, syntax 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.

1 participant