Skip to content

Chapter 12

Overview

Flax is JAX's ML framework, providing modules, state management, and best practices.

Key Topics

1. Basic Module

import flax.linen as nn
import jax.numpy as jnp

class MLP(nn.Module):
 hidden_size: int
 output_size: int

 @nn.compact
 def __call__(self, x):
 x = nn.Dense(self.hidden_size)(x)
 x = nn.relu(x)
 x = nn.Dense(self.output_size)(x)
 return x

# Create and initialize
model = MLP(hidden_size=256, output_size=10)
params = model.init(jax.random.PRNGKey(0), jnp.ones((1, 784)))

# Forward pass
output = model.apply(params, x)

2. Training State

from flax.training import train_state
import optax

# Create optimizer
optimizer = optax.adam(learning_rate=1e-3)

# Create train state
state = train_state.TrainState.create(
 apply_fn=model.apply,
 params=params,
 tx=optimizer
)

3. Training Step

def loss_fn(params, x, y):
 logits = model.apply(params, x)
 loss = jnp.mean((logits - y)**2)
 return loss

def train_step(state, x, y):
 grads = grad(loss_fn)(state.params, x, y)
 return state.apply_gradients(grads=grads)

# Training loop
for epoch in range(100):
 state = train_step(state, x_batch, y_batch)

Advantages

  • Cleaner API than manual JAX
  • pytrees automatically handled
  • Common patterns (dropout, batch norm)
  • Production-ready

Summary

  • Modules as JAX best practice
  • TrainState for parameter management
  • Integrates with optax optimizers
  • Recommended for production JAX code