Skip to content

Chapter 13

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