Skip to content

Treat integer hyperparameters as static in inject_hyperparams (fixes #412) - #1730

Open
ArneshBanerjee wants to merge 2 commits into
google-deepmind:mainfrom
ArneshBanerjee:fix-inject-hyperparams-integer-static
Open

ArneshBanerjee wants to merge 2 commits into
google-deepmind:mainfrom
ArneshBanerjee:fix-inject-hyperparams-integer-static

Conversation

@ArneshBanerjee

Copy link
Copy Markdown

Summary

optax.inject_hyperparams converts every numeric argument into a traced array so it can be scheduled/overridden at runtime. Boolean arguments are already special-cased as static, because a traced boolean can't be used in Python control flow.

Integer arguments have the identical problem, but currently fall through to the "numeric" branch and get traced. Several optimizers use integer arguments for structural decisions:

  • min_dim_size_to_factor in adafactor (optax/_src/factorized.py:55 → if shape[...] < min_dim_size_to_factor)
  • memory_size in lbfgs (optax/_src/transform.py:1717 → if memory_size < 1)

As a result, jitting these fails with TracerBoolConversionError unless the user manually passes static_args=(...):

import jax, jax.numpy as jnp, optax
opt = optax.inject_hyperparams(optax.adafactor)(learning_rate=0.1)
jax.jit(opt.init)(jnp.ones((4, 4)))   # TracerBoolConversionError

This is exactly what #412 asks to fix ("all optimizers wrapped in inject_hyperparams can be jit compiled without any additional static_args").

Fix

Since bool is a subclass of int, extending the existing static-value check from bool to int generalizes the current behavior and resolves both optimizers with no new special cases. Integer hyperparameters cannot be meaningfully scheduled under jit anyway (a schedule returns floats, and structural ints break control flow when traced).

Verification

  • inject_hyperparams(optax.adafactor) and inject_hyperparams(optax.lbfgs) now jax.jit out of the box; the previously-required static_args workaround is no longer needed.
  • Added regression tests in optax/schedules/_inject_test.py.
  • pytest optax/schedules/ optax/_src/alias_test.py → 634 passed, 70 skipped, no regressions.

Fixes #412.

@google-cla

google-cla Bot commented Jul 21, 2026

Copy link
Copy Markdown

Thanks for your pull request! It looks like this may be your first contribution to a Google open source project. Before we can look at your pull request, you'll need to sign a Contributor License Agreement (CLA).

View this failed invocation of the CLA check for more information.

For the most up to date status, view the checks section at the bottom of the pull request.

@ArneshBanerjee
ArneshBanerjee force-pushed the fix-inject-hyperparams-integer-static branch from bba7c36 to d6eb26b Compare July 21, 2026 20:05
@ArneshBanerjee
ArneshBanerjee force-pushed the fix-inject-hyperparams-integer-static branch from d6eb26b to 8f973c3 Compare August 18, 2026 16:20
@ArneshBanerjee

Copy link
Copy Markdown
Author

Rebased on main; CLA is signed and CI is green.

This fixes #412: inject_hyperparams traced integer hyperparameters, which breaks inner factories that use them for structural control flow (min_dim_size_to_factor in adafactor, memory_size in lbfgs) with a TracerBoolConversionError under jit. Integers are now treated as static, matching how bools were already handled.

Would someone be able to review?

@ArneshBanerjee
ArneshBanerjee force-pushed the fix-inject-hyperparams-integer-static branch from 8f973c3 to 440e065 Compare September 3, 2026 08:10

@sylvesterkaczmarek sylvesterkaczmarek left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The new rule only recognizes Python int; zero-dimensional integer NumPy/JAX arrays still fall into numeric_hps and get traced. Thus a structural argument such as memory_size=jnp.asarray(10) can still hit the same TracerBoolConversionError this change is intended to prevent. Please classify integer-dtype scalar arrays as static too, or key the rule to structural parameters, and test NumPy/JAX integer scalars.

@ArneshBanerjee

ArneshBanerjee commented Sep 3, 2026 •

Copy link
Copy Markdown
Author

I reproduced it: with memory_size=jnp.asarray(10) or np.asarray(10) the value still landed in numeric_hps and got traced, so it hit the same TracerBoolConversionError. (np.int64(10) happened to work already, but only by accident, since it is not an ndarray and fell through to the static branch.)

Pushed a fix:

  1. The static rule now matches zero-dimensional NumPy and JAX values of integer or boolean dtype, not just Python int.

  2. Classifying them as static is not sufficient on its own for jax.Array. An array closed over by a jitted function is staged out as a constant of the trace, so using it for structural control flow still raised the same error even once it was no longer injected. Static integer and boolean values are now converted to Python scalars so they are real compile-time constants. A traced value is passed through unchanged so the inner factory reports it rather than this raising something less clear.

I limited the rule to zero-dimensional values on purpose. A raw jax.random.PRNGKey is uint32 with shape (2,), so anything keyed on integer dtype alone would have made keys static too. There is a test covering both key types to pin that down.

Tests added over int, bool, np.int64, np.asarray, jnp.asarray and a bool array, checking the hyperparameter stays static and that the update jits. They fail on the previous commit for the three array cases.

@sylvesterkaczmarek sylvesterkaczmarek left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Zero-dimensional NumPy/JAX integer and boolean values are now classified as static and converted to Python scalars before they can be closed over by jit as traced arrays. The new tests cover the structural integer cases that previously still raised TracerBoolConversionError. My finding is resolved.

@selamw1
selamw1 self-requested a review October 1, 2026 21:31

@selamw1 selamw1 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for the PR! I tested it locally and it fixes #412:

  • jax.jit(inject_hyperparams(adafactor)(learning_rate=0.1).init) now works, as do inject_stateful_hyperparams and lbfgs.
  • 7 of the 11 new tests fail without the fix and pass with it.
  • alias_test.py, schedules/*, transforms/* and contrib/_common_test.py give the same results before and after.

A few concerns before merging:

1. Breaking change: integer values for float hyperparameters are no longer injected

The rule checks the type of the value passed, not what the argument means. So learning_rate=1, b1=0 or momentum=0 (and, since 39bc54d, jnp.asarray(1)) become static and silently disappear from state.hyperparams:

opt = optax.inject_hyperparams(optax.sgd)(learning_rate=1)
state = opt.init(params)
'learning_rate' in state.hyperparams   # before: True, after: False
state.hyperparams['learning_rate'] = jnp.asarray(0.5)
opt.update(grads, state)
# after: TypeError: sgd() got multiple values for keyword argument 'learning_rate'

1.0 still works, so this is easy to miss. A safer rule would use the factory's signature: keep an argument static only if its default is an int or its annotation is int/bool. That covers min_dim_size_to_factor, memory_size, rank, update_proj_gap, num_betas, warmup_steps and muon's ns_steps, without catching float arguments.

If you keep the current approach, please add a test for this case. Please also update the static_args docstring: it currently says integers "are always treated as static" without mentioning that this includes integers passed for float arguments.

2. Remaining items from #412

  • Remove the existing static_args workarounds in the tests. alias_test.py still passes static_args=('min_dim_size_to_factor',) for adafactor (with a comment pointing to #412). contrib/_common_test.py still passes ['warmup_steps', 'num_betas', 'clip_value_fn', 'ns_steps', 'rank', 'update_proj_gap']. With this PR applied, I removed them locally and all of those tests pass except adopt. Removing them would show the fix works and guard against regressions.
  • contrib.adopt still fails under jit ('DynamicJaxprTracer' object is not callable) unless static_args=("clip_value_fn",) is passed, because its default clip_value_fn is a function and gets treated as a schedule. Since #412 suggests no optimizer should need static_args, it would be good to handle this here or in a follow-up.

3. Unrelated file changes

The diff seems to include unrelated files. This PR should only touch optax/schedules/_inject.py, optax/schedules/_inject_test.py, and the test cleanups above.

Happy to re-test once these are addressed.

`inject_hyperparams` converts every numeric argument into a traced array so
it can be scheduled. Boolean arguments were already special-cased as static,
because a traced boolean breaks Python control flow. Integer arguments have
the exact same problem: several optimizers use them for structural decisions,
e.g. `min_dim_size_to_factor` in `adafactor` (factorized.py) and `memory_size`
in `lbfgs`. Injecting them as traced arrays raises a
`TracerBoolConversionError` when the resulting transform is jitted, forcing
users to manually pass `static_args=(...)`.

Since `bool` is a subclass of `int`, treating `int` as static generalizes the
existing behavior and lets `inject_hyperparams(optax.adafactor)` and
`inject_hyperparams(optax.lbfgs)` be jitted out of the box. Integer
hyperparameters cannot be meaningfully scheduled under jit anyway.

Fixes google-deepmind#412.
Integer arguments declared as int or bool (by annotation or default) are
now kept out of the injected hyperparameters, so factories that use them
for control flow, such as adafactor and lbfgs, can be jitted without
static_args. Zero-dimensional integer arrays are converted to Python
scalars. Integers passed for float arguments, like learning_rate=1, are
still injected.

Remove the static_args workarounds from the alias and contrib tests,
except clip_value_fn for adopt.

Fixes google-deepmind#412.
@ArneshBanerjee
ArneshBanerjee force-pushed the fix-inject-hyperparams-integer-static branch from 39bc54d to 8347148 Compare October 2, 2026 13:00
@ArneshBanerjee

Copy link
Copy Markdown
Author

Thanks for testing this, all three points were right.

  1. The rule now uses the factory's signature. An integer value (or a 0-d integer array) is static only if the parameter is annotated int/bool or has an int default. learning_rate=1 is injected again, and there is a test for it. The static_args docstring says this now.
  2. Removed the static_args workarounds in alias_test.py and contrib/_common_test.py. Only clip_value_fn stays for adopt, which I would rather fix in a follow-up.
  3. The extra files came from a file mode change (644 to 755) in my second commit. I squashed everything into one commit on current main, and it now touches 4 files.

@sylvesterkaczmarek

Copy link
Copy Markdown

Re-reviewed current head 8347148. The signature-based rule addresses the compatibility concern from the latest review: structural integer/bool parameters stay static, while an integer value passed to a float parameter such as learning_rate=1 remains injected. The old adafactor/contrib static-arg workarounds are removed, the scalar NumPy/JAX cases remain covered, and the optimizer/JAX test matrix is green. The remaining markdown-link failure is from GitHub URLs returning 503 and is unrelated to this patch. No remaining blocker from my review.

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.

Problems when jitting Adafactor with inject_hyperparams.

3 participants