This guide demonstrates how to perform data-parallel distributed training using Keras 3 with the JAX backend on Google Cloud TPU VMs. It covers setting up the environment, building a model, configuring JAX sharding for data and variables, and implementing a custom training loop using Keras' stateless APIs.
Prerequisites
- A Google Cloud TPU VM with at least 8 local devices.
- Follow the official instructions to spin up a TPU VM: https://cloud.google.com/tpu/docs/run-calculation-jax
Setup and Configuration
- Force the JAX backend by setting the environment variable before importing Keras:
import os
os.environ["KERAS_BACKEND"] = "jax"
- Import necessary libraries including
jax, jax.sharding, and keras.
Model and Data Preparation
- Build a Keras
Sequential or functional model (e.g., a CNN for MNIST). - Load data using
tf.data (compatible with Keras) and convert to batches. - Initialize the model and optimizer state using
.build() with a dummy batch:
model.build(one_batch)
optimizer.build(model.trainable_variables)
JAX Distribution Setup
- Create a JAX device mesh and sharding configurations:
- Data Sharding: Split data along the batch axis (
P("batch")). - Variable Replication: Replicate variables across all devices (
P()). - Custom Sharding: Optionally shard specific large kernels (e.g., split a Conv2D kernel across 4 devices).
Example mesh and sharding setup:
from jax.experimental import mesh_utils
from jax.sharding import Mesh, NamedSharding, PartitionSpec as P
devices = mesh_utils.create_device_mesh((8,))
data_mesh = Mesh(devices, axis_names=("batch"))
data_sharding = NamedSharding(data_mesh, P("batch"))
var_mesh = Mesh(devices, axis_names=("_"))
var_replication = NamedSharding(var_mesh, P())
Custom Training Loop with Stateless APIs
- Use
model.stateless_call for the forward pass and optimizer.stateless_apply for updates. These functions are backend-agnostic and required for functional JAX workflows. - Define a loss function and a gradient computation function using
jax.value_and_grad. - JIT-compile the training step:
@jax.jit
def train_step(train_state, x, y):
(loss_value, non_trainable_variables), grads = compute_gradients(
train_state.trainable_variables,
train_state.non_trainable_variables,
x,
y,
)
trainable_variables, optimizer_variables = optimizer.stateless_apply(
train_state.optimizer_variables, grads, train_state.trainable_variables
)
return loss_value, TrainingState(
trainable_variables, non_trainable_variables, optimizer_variables
)
Inference and State Synchronization
- Run predictions using
model.stateless_call with the sharded data. - After training, update the original Keras model variables using
jax.tree_map and variable.assign to ensure the model state is synchronized for subsequent evaluation or saving:
update = lambda variable, value: variable.assign(value)
jax.tree_map(update, model.trainable_variables, device_train_state.trainable_variables)
jax.tree_map(update, model.non_trainable_variables, device_train_state.non_trainable_variables)
- Compile the model and run
model.evaluate() to verify the updated state.
import os
os.environ["KERAS_BACKEND"] = "jax"
import jax
import jax.numpy as jnp
import keras
from jax.experimental import mesh_utils
from jax.sharding import Mesh, NamedSharding, PartitionSpec as P
# 1. Setup Mesh and Sharding
devices = mesh_utils.create_device_mesh((8,))
data_mesh = Mesh(devices, axis_names=("batch"))
data_sharding = NamedSharding(data_mesh, P("batch"))
var_mesh = Mesh(devices, axis_names=("_"))
var_replication = NamedSharding(var_mesh, P())
# 2. Build and Build State
model = keras.Sequential([...]) # Define model
optimizer = keras.optimizers.Adam(0.01)
model.build(dummy_batch)
optimizer.build(model.trainable_variables)
# 3. Sharding Variables
trainable_variables = [jax.device_put(v, var_replication) for v in model.trainable_variables]
# ... custom sharding logic for specific layers ...
# 4. Stateless Training Step
@jax.jit
def train_step(state, x, y):
(loss, non_trainable), grads = jax.value_and_grad(
lambda tv, ntv, x, y: keras.losses.SparseCategoricalCrossentropy()(y, model.stateless_call(tv, ntv, x)[0]),
has_aux=True
)(state.trainable_variables, state.non_trainable_variables, x, y)
tv, opt_vars = optimizer.stateless_apply(state.optimizer_variables, grads, state.trainable_variables)
return loss, TrainingState(tv, non_trainable, opt_vars)
# 5. Sync State Back to Model
jax.tree_map(lambda v, val: v.assign(val), model.trainable_variables, new_state.trainable_variables)
Sources: examples/demo_jax_distributed.py