Skip to content

Give training shared variables a static shape - #60

Open
bwengals wants to merge 1 commit into
mainfrom
fix/static-shared-shapes
Open

bwengals wants to merge 1 commit into
mainfrom
fix/static-shared-shapes

Conversation

@bwengals

@bwengals bwengals commented Sep 29, 2026 •

Copy link
Copy Markdown
Collaborator

Closes #54.

_make_shared_params created the trainable shared slots without shape=, so an RV declared with dims and no shape kept a type shape of (None, None) after replacement. Under JAX (and MLX), an arange whose bound comes from such a shape is not concrete, and a trace penalty like pt.trace(W @ W.T) lowers to one in the gradient, so compilation failed.

The shared slots now take their shape from the init value, as proposed in the issue. Parameters do not resize during a fit, and every existing set_value path (unpack_to_shared, load_fit, staged VFE) writes a value of the same shape, so nothing that worked before is affected. The frozen-Z fallback in minimize_staged_vfe gets the same treatment.

Tests

  • test_shared_params_have_static_shape: backend independent. Shared slots for dims-only RVs have concrete shapes (W: (4, 2), kappa_log__: (4,)).
  • test_dims_only_rv_compiles_under_jax: the issue's reproduction, compiled with mode="JAX". Skipped when JAX is not installed, which includes CI since JAX is not a test dependency.

Both fail on main with the reported errors and pass with this change. The full suite passes locally (285 tests).


📚 Documentation preview 📚: https://ptgp--60.org.readthedocs.build/en/60/

This branch has not been deployed

No deployments
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.

Give compile_training_step's shared parameters a static shape for JAX and MLX

1 participant