Treat integer hyperparameters as static in inject_hyperparams (fixes #412) - #1730
ArneshBanerjee wants to merge 2 commits into
Conversation
|
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. |
bba7c36 to
d6eb26b
Compare
d6eb26b to
8f973c3
Compare
|
Rebased on main; CLA is signed and CI is green. This fixes #412: Would someone be able to review? |
8f973c3 to
440e065
Compare
sylvesterkaczmarek
left a comment
There was a problem hiding this comment.
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.
|
I reproduced it: with Pushed a fix:
I limited the rule to zero-dimensional values on purpose. A raw Tests added over |
sylvesterkaczmarek
left a comment
There was a problem hiding this comment.
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
left a comment
There was a problem hiding this comment.
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 doinject_stateful_hyperparamsandlbfgs.- 7 of the 11 new tests fail without the fix and pass with it.
alias_test.py,schedules/*,transforms/*andcontrib/_common_test.pygive 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_argsworkarounds in the tests.alias_test.pystill passesstatic_args=('min_dim_size_to_factor',)for adafactor (with a comment pointing to #412).contrib/_common_test.pystill 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 exceptadopt. Removing them would show the fix works and guard against regressions. contrib.adoptstill fails under jit ('DynamicJaxprTracer' object is not callable) unlessstatic_args=("clip_value_fn",)is passed, because its defaultclip_value_fnis a function and gets treated as a schedule. Since #412 suggests no optimizer should needstatic_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.
39bc54d to
8347148
Compare
|
Thanks for testing this, all three points were right.
|
|
Re-reviewed current head |
Summary
optax.inject_hyperparamsconverts 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_factorinadafactor(optax/_src/factorized.py:55→if shape[...] < min_dim_size_to_factor)memory_sizeinlbfgs(optax/_src/transform.py:1717→if memory_size < 1)As a result, jitting these fails with
TracerBoolConversionErrorunless the user manually passesstatic_args=(...):This is exactly what #412 asks to fix ("all optimizers wrapped in
inject_hyperparamscan be jit compiled without any additionalstatic_args").Fix
Since
boolis a subclass ofint, extending the existing static-value check frombooltointgeneralizes 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)andinject_hyperparams(optax.lbfgs)nowjax.jitout of the box; the previously-requiredstatic_argsworkaround is no longer needed.optax/schedules/_inject_test.py.pytest optax/schedules/ optax/_src/alias_test.py→ 634 passed, 70 skipped, no regressions.Fixes #412.