Skip to content

Preserve plateau scale dtype with array minimum scales - #1774

Open
sylvesterkaczmarek wants to merge 2 commits into
google-deepmind:mainfrom
sylvesterkaczmarek:fix/plateau-array-min-scale-dtype
Open

sylvesterkaczmarek wants to merge 2 commits into
google-deepmind:mainfrom
sylvesterkaczmarek:fix/plateau-array-min-scale-dtype

Conversation

@sylvesterkaczmarek

Copy link
Copy Markdown

Array-like min_scale values can promote the reduced scale above the dtype stored in ReduceLROnPlateauState. The cooldown branches then return different dtypes, so lax.cond raises a TypeError on the first update with float16 or bfloat16 parameters and a float32 minimum scale.

Cast the result back to the existing scale dtype. Add coverage for NumPy/JAX minimum scales, mixed-precision parameter trees, accumulation, cooldown, and repeated reductions down to the floor, both eagerly and under JIT.

Validation

  • Scheduler and hyperparameter-injection tests: 53 passed on CPU.
  • The new cases on unpatched main: 20 failed, 4 passed.
  • Pre-commit and Pylint pass.

This is separate from the negative-metric comparison change in #1767.

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