google / google/flax

Migrate vae example from flax.linen to flax.nnx

Open
#5,068 2 comments 0 reactions 0 assignees View on GitHub
Dominant language
Jupyter Notebook
Stars
7.3k
Forks
833
Avg merge
5h 11m
Merged PRs (30d)
5

Description

#### Description
The current [VAE example](https://github.com/google/flax/tree/main/examples/vae) uses the `flax.linen` API for model definition and training.
As Flax continues to develop the `nnx` module as its next-generation neural network API, it would be valuable to provide an updated version of this example using `flax.nnx`.

This migration will help users:
- Learn how to implement a VAE using `nnx`'s new modular and explicit state-handling paradigm.
- Compare differences between `nn` and `nnx` APIs in real-world use cases.
- Encourage adoption of `nnx` in research and production examples.

---

#### Proposed Changes
- Reimplement the model (`Encoder`, `Decoder`, and `VAE` wrapper) using `flax.nnx.Module`.
- Replace `flax.training.train_state.TrainState` with `nnx.Optimizer` for parameter management.
- Update training and evaluation loops to use `nnx.jit` and direct method calls instead of `apply()`.
- Ensure reproducibility and equivalence with the original `nn`-based example.

---

#### Contribution
I'd be happy to implement this migration and submit a PR. Please let me know if there are any specific guidelines or preferences for the implementation approach.

---

#### Motivation
The VAE example is a widely understood benchmark that involves both deterministic and stochastic components, making it ideal to showcase `nnx`'s design strengths:
- Explicit randomness (`nnx.Rngs`)
- Parameter/state separation
- Compositional design
- Compatibility with Optax and other JAX tools

Having this example available in `nnx` would significantly benefit users exploring or transitioning to the new API.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.