Chapter 13: Training Loops & Advanced Optimization¶
Overview¶
Building efficient training loops with JAX: gradient accumulation, mixed precision, learning rate schedules.
Topics¶
1. Gradient Accumulation¶
import jax
import jax.numpy as jnp
def accumulate_gradients(state, batch_size, x, y):
accum_grads = None
for i in range(0, len(x), batch_size):
x_batch = x[i:i+batch_size]
y_batch = y[i:i+batch_size]
grads = grad(loss_fn)(state.params, x_batch, y_batch)
if accum_grads is None:
accum_grads = grads
else:
accum_grads = jax.tree_map(
lambda g1, g2: g1 + g2,
accum_grads, grads
)
return accum_grads
2. Learning Rate Schedule¶
import optax
# Exponential decay
schedule = optax.exponential_decay(
init_value=0.1,
transition_steps=1000,
decay_rate=0.96
)
# Cosine annealing
schedule = optax.cosine_decay_schedule(
init_value=0.1,
decay_steps=10000
)
# Use with optimizer
optimizer = optax.chain(
optax.clip_by_global_norm(1.0), # Gradient clipping
optax.adam(learning_rate=schedule)
)
3. Mixed Precision Training¶
import jax.numpy as jnp
from jax import dtype as jdt
def loss_fn_mixed(params, x, y):
# Compute in float32
logits = model.apply(params, x.astype(jnp.float32))
# Loss in float32
loss = jnp.mean((logits - y)**2)
return loss
# Params in float16, but JAX handles precision
params_fp16 = jax.tree_map(lambda x: x.astype(jnp.float16), params)
4. Multi-GPU Training¶
from jax.experimental import multihost_utils
# Use pmap for data parallelism
def train_step_pmap(params_and_state, batch):
# This runs on each device
state, params = params_and_state
loss, grads = compute_loss_and_grad(params, batch)
return (state, params), loss
# Replicate across devices
n_devices = jax.device_count()
batches = split_batch(x_train, n_devices)
# pmap handles distribution
train_step_vmapped = jax.pmap(train_step_pmap)
Summary¶
- Gradient accumulation for large batches
- Learning rate schedules for convergence
- Mixed precision for memory/speed
- pmap for multi-GPU training