Library / SDK
google/flax avatar
google/flax

Flax: the NNX API for JAX, and what Linen users should know

Flax is a neural network library for JAX that is designed for flexibility.

7,333 stars846 forksJupyter NotebookApache-2.0

At a glance

What is it?
Flax is Google's neural network library for JAX. Released in 2024, NNX replaces Linen's functional module style with plain Python objects, and the repository now carries two APIs with two documentation sites.
Who is it for?
Adopt Flax if you already write JAX and want nnx.Module to behave like a normal Python object, with the MNIST tutorial and the Gemma example as your starting points. Do not adopt it if you want a framework that owns the training loop, or if you need Linen support to keep growing rather than being maintained for existing code.
Can I use it commercially?
Yes. Apache-2.0 is a permissive licence: you can use, modify and sell software built on it, as long as you keep its copyright and licence notices.
Is it still maintained?
Yes. The repository last received commits 4 days ago.
What is it written in?
Mainly Jupyter Notebook, according to GitHub's language statistics.

Answers come from the project's GitHub data, last synced on September 29, 2026, and from our analysis. They are not legal advice.

Editorial analysis

The problem Flax solves for people already writing JAX

JAX gives you transformations (jit, grad, vmap) and arrays, not a model abstraction. Flax is the layer above that. The README describes it as a neural network library and ecosystem for JAX that is designed for flexibility, and the stated design intent is that you try new forms of training by forking an example and modifying the training loop, rather than asking the framework to add a feature. That is a deliberate position: Flax does not want to own your training loop.

The audience follows from that. The Flax team's stated mission is to serve the JAX neural network research ecosystem, both inside Alphabet and in the wider community. If you are writing a training loop by hand and only want parameter containers, layers and checkpointing, Flax fits. If you want a framework that decides the loop for you, this is not it.

The API story matters more than the tagline. Linen was released in 2020 by engineers and researchers at Google Brain in collaboration with the JAX team. NNX was released in 2024 as a simplified API that adds first class support for Python reference semantics, so models are expressed as regular Python objects with reference sharing and mutability. Two APIs, two documentation sites, one repository.

How NNX modules actually work: reference semantics instead of functional purity

The README's MLP example is the clearest statement of the mechanism. A model subclasses nnx.Module, and its submodules are created in __init__ as attributes: nnx.Linear, nnx.Dropout, nnx.BatchNorm. The forward pass is __call__, and it reads like ordinary Python. It calls nnx.gelu, the dropout layer, the batch norm layer and the second linear layer in sequence.

That is the whole point of reference semantics. Layers are objects held on the module, not values threaded through a functional transform, and they can be shared and mutated. The README says this design is meant to make networks easier to create, inspect, debug and analyze. Inspection is the practical payoff: when a model is a Python object graph, a debugger and a repr can show you what is inside it.

Randomness is explicit. Constructors take an rngs argument, typed as nnx.Rngs. In the MLP every layer that needs randomness receives the same rngs object, and Dropout takes both rate and rngs. This is how NNX keeps JAX's explicit PRNG handling without making you pass keys through every call.

What ships alongside the layers is the rest of the list in the README: replicated training, serialization and checkpointing, metrics, and prefetching on device. Checkpointing depends on orbax-checkpoint and tensorstore, which appear as install dependencies in pyproject.toml. The layer set is broad enough to be useful without leaving the library: Linear, Conv, BatchNorm, LayerNorm, GroupNorm, MultiHeadAttention, LSTMCell, GRUCell and Dropout are all named in the README's neural network API list.

Installing Flax and defining a first model

Flax uses JAX, so the README points you at the JAX installation instructions for CPUs, GPUs and TPUs before anything else. Get JAX working for your accelerator first, then install Flax from PyPI:

bash
pip install flax

If you need optional dependencies that some dependencies do not pull in, such as matplotlib, the README gives an extras install:

bash
pip install "flax[all]"

To move to the latest code rather than the released version, the README documents an upgrade straight from the repository:

bash
pip install --upgrade git+https://github.com/google/flax.git

One version constraint to check before you start: pyproject.toml sets requires-python to >=3.12, and the classifiers list Python 3.12. The README's quick install section still says you need Python 3.8 or later, which contradicts the packaging metadata. Trust the metadata when you provision an environment.

With Flax installed, the README's MLP is a complete first model. Save it as a module and instantiate it with an Rngs object:

python
import jax
from flax import nnx

class MLP(nnx.Module):
  def __init__(self, din: int, dmid: int, dout: int, *, rngs: nnx.Rngs):
    self.linear1 = nnx.Linear(din, dmid, rngs=rngs)
    self.dropout = nnx.Dropout(rate=0.1, rngs=rngs)
    self.bn = nnx.BatchNorm(dmid, rngs=rngs)
    self.linear2 = nnx.Linear(dmid, dout, rngs=rngs)

  def __call__(self, x: jax.Array):
    x = nnx.gelu(self.dropout(self.bn(self.linear1(x))))
    return self.linear2(x)

What you should see is an object you can print and walk: the four submodules are attributes on the instance, and calling the model with an array returns the output of linear2. The README does not document a training loop for this snippet, so pair it with your own loop or with the MNIST tutorial linked from the README.

The Linen question, and why the repository carries two documentation sites

The most consequential thing about Flax right now is that it is two libraries in one repository. Linen is the 2020 API. NNX is the 2024 API. The README links NNX documentation, and it states plainly that Flax Linen's documentation has its own site. It also links a guide on the evolution from Linen to NNX.

That split has a cost the README does not hide but does not dwell on either. Search results and tutorials written before 2024 describe Linen modules, functional apply calls and parameter dictionaries. NNX code looks nothing like that. A reader who follows an older tutorial and then installs the current release will be working against a different mental model, and the repository layout reflects the duplication: there is both a docs/ directory and a docs_nnx/ directory, and the examples tree contains a linen_design_test directory alongside nnx_toy_examples.

On stability, the README says the team expects to improve Flax but does not anticipate significant breaking changes to the core API, and that it uses changelog entries and deprecation warnings where possible. That is a reasonable promise for a research library, and it is weaker than a compatibility guarantee. Note also that pyproject.toml classifies the package as Development Status :: 3 - Alpha, which sits oddly next to the deprecation-warning policy. Treat the classifier as the more conservative signal when you plan upgrades.

Where Flax is the wrong tool

If you want a framework to own the training loop, Flax is pointed the other way. The README's own framing is that flexibility comes from forking an example and modifying the loop, not from adding features to a framework. Teams that want callbacks, a fit method and a fixed project structure will spend their first week rebuilding what a higher-level library gives them.

JAX is a hard prerequisite, not an optional backend. The README sends you to the JAX installation instructions before the install command, and pyproject.toml pins jax>=0.11.1 with numpy, optax, orbax-checkpoint, tensorstore, msgpack, rich, typing_extensions, PyYAML and treescope. If your environment is not already a JAX environment, the dependency set is the first thing you will be managing.

The Python floor is another boundary. requires-python is >=3.12, so any pipeline still on 3.10 or 3.11 cannot install the current release without changing the interpreter. And the optional testing extras show how heavy the ecosystem around the examples is: tensorflow, tensorflow_text, tensorflow_datasets, jaxlib, jraph, gymnasium and keras all appear there. The core install is small. Reproducing the repository's own examples is not.

Finally, the two-API situation is a real adoption risk. Starting new code in Linen in 2026 means starting on the API that NNX evolved from, with its own separate documentation site, while the README's attention is on NNX.

Flax against PyTorch: the difference is who owns the loop

The comparison people search for is Flax versus PyTorch, and the difference is structural rather than a matter of layer coverage. Both give you Linear, Conv, normalization, attention and recurrent cells. The divergence is what happens around them.

PyTorch's eager execution and autograd make the model a mutable object by default, and the surrounding ecosystem supplies training loops, data loading and deployment paths. Flax gets to mutability through NNX's reference semantics, which the README describes as first class support for Python reference semantics, but it reaches that point on top of JAX's functional transformations. The library does not hand you a training loop, and the README says that is intentional.

A second difference is scale-out. Flax's stated collaboration is with the JAX team, and the README's utility list includes replicated training and prefetching on device. The examples tree is the evidence for what the team considers representative: imagenet, lm1b, wmt, seq2seq, ppo, ogbg_molpcba, sst2 and a Gemma inference example. That is a research and large-model set, not a web-application set.

If your team already knows PyTorch and has no JAX investment, Flax asks you to learn JAX's transformation model before it gives you anything. If your team is already writing JAX, the comparison inverts: Flax is the model layer you would otherwise write yourself.

Maintenance, releases and the Apache-2.0 licence

The repository is not archived, and the last push was on 2026-09-16. Releases have been coming at a steady cadence: v0.12.7 on 2026-04-22, v0.12.8 on 2026-07-20, and v0.12.9 on 2026-08-18. Versioning is dynamic, driven by setuptools-scm, so the version comes from tags rather than a literal in pyproject.toml.

Upgrade cost is mostly the API split. The README says the team uses changelog entries and deprecation warnings where possible, and the repository has a CHANGELOG.md at the top level. The practical upgrade path is to read that file per release and to check the Linen to NNX guide if you are still on Linen. Because the package is classified as Alpha, pinning a version in production is the safer default, and the README's git-based upgrade command is for tracking main rather than for a pinned deployment.

The licence is Apache-2.0, declared in the LICENSE file and in the classifiers as License :: OSI Approved :: Apache Software License. That is a permissive licence, and it is the same licence family used across much of the JAX ecosystem, which keeps dependency licence review simple. This is a description of what the repository declares, not legal advice; if your organisation has a licence policy, run the LICENSE file through it.

Editorial conclusion

Adopt Flax if you already write JAX and want nnx.Module to behave like a normal Python object, with the MNIST tutorial and the Gemma example as your starting points. Do not adopt it if you want a framework that owns the training loop, or if you need Linen support to keep growing rather than being maintained for existing code. Before committing, check that your Python version satisfies requires-python >=3.12 in pyproject.toml, and read the Linen to NNX guide to see how much of your existing code has to change.

Frequently asked questions

What is Flax used for?

Flax is a neural network library and ecosystem for JAX. It provides a neural network API in flax.nnx along with utilities for replicated training, serialization and checkpointing, metrics and prefetching on device, plus educational examples such as MNIST and Gemma inference.

What is Flax and JAX?

JAX is the array and transformation library that Flax is built on; the README sends readers to the JAX installation instructions before installing Flax. Flax adds the model layer on top, and it is developed in close collaboration with the JAX team.

What is the difference between Flax NNX and Flax Linen?

Linen was released in 2020 by engineers and researchers at Google Brain. NNX was released in 2024 as a simplified API that adds first class support for Python reference semantics, so models are regular Python objects with reference sharing and mutability. Linen keeps its own documentation site.

How do I install Flax?

Install JAX for your CPU, GPU or TPU first, following the JAX installation instructions the README links to, then run pip install flax. The README also documents pip install "flax[all]" for extra optional dependencies.

Which Python version does Flax require?

The packaging metadata in pyproject.toml sets requires-python to >=3.12 and lists Python 3.12 in its classifiers. The README's quick install section says Python 3.8 or later, which does not match the metadata, so check pyproject.toml when provisioning an environment.

Official sources

  1. google/flax on GitHub
  2. License: Apache-2.0
  3. Project website
  4. README
  5. Releases
Add this badge to your README

If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.

Add this badge to your README

markdown
[![Hysen Labs](https://hysenlabs.com/badge/google-flax.svg)](https://hysenlabs.com/projects/google-flax)