Skip to content

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