Chapter 12: Flax Framework¶
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