Repository navigation
Add more examples in the NNX docs #5433
Description
Activity
Hi @vfdev-5, I would like to pick up the "Time series classification with CNN" example if nobody is already working on it.
I plan to add a small NNX-based tutorial/example covering:
- a simple 1D CNN model for time-series classification
- synthetic or lightweight dataset setup to keep the example easy to run
- training/eval loop using JAX + Flax NNX
- metrics/loss reporting
- docs style consistent with the existing NNX examples
I’ll first check the existing examples structure and docs build setup, then open a draft PR for feedback before expanding it too much.
@cgarciae @samanklesaria do we want to provide an example on
- "Text classification with a transformer language model using JAX" ?
- "Time series classification with CNN" ?
Hi @vfdev-5 , I'd like to take this on. I've built neural nets before and I'm learning the NNX API, so I think I can write an example that's genuinely clear for newcomers.
My proposal: a self-contained, well-commented tutorial that walks through a full training pipeline in NNX end-to-end — defining a Module, the training step with nnx.jit / nnx.value_and_grad, the optimizer, and an eval loop — framed for people migrating from PyTorch/Keras. Since MLP/CNN/autoencoder already exist, I'd focus on making the workflow and idioms the teaching point rather than the architecture.
Does that fit what you had in mind, or would you prefer I target a different gap? Once you confirm the direction I'll open a PR following the contributing guidelines, tested against the current NNX API.
Hi @vfdev-5 — I'd like to pick up the unchecked "Part 2: Debug a variational autoencoder (VAE)" item if it's still open.
I've gone through the source tutorial (https://docs.jaxstack.ai/en/latest/digits_vae.html). It already uses
flax.nnx—nnx.Module,nnx.Linear— so my read is that the work here is moving it into the NNX docs, refreshing it against the current NNX API, and wiring it into the toctree, following the same shape as the earlier items on this list rather than a rewrite. It uses the scikit-learn digits dataset, so it should stay cheap to run in the docs build.I'll start on it and open a draft PR so there's something concrete to look at. Happy to adjust the scope once it's up if you had something broader in mind.
We plan for moving Jax AI Stack examples to NNX docs:
Image Captioning with Vision Transformer (ViT) model (@vfdev-5, ...)