A TPU is not something you program kernel by kernel. You hand a whole training step to a compiler, the compiler produces one program for every chip in the slice, and from then on the chips run that program in lock step while the hosts attached to them do nothing but feed data and save checkpoints. That division of labour explains almost every success and every failure you will see on TPU v5. When a v5 job is fast, it is because the compiled program kept the matrix units busy and the hosts kept up. When it is slow, it is usually because a new input shape forced a recompile, a host could not decode data quickly enough, or a collective that should have stayed on the chip-to-chip interconnect went over the data-centre network instead.

This article is about running training on the two TPU v5 chips, v5e and v5p, from the software side: how the hardware is named and provisioned, what each host process does, how to write a step that compiles once, how to size the input pipeline with a worked example, how Multislice spans slices, and how to survive maintenance and preemption. The specification comparison, roofline arithmetic and generation choice are covered in TPU v5 and v6 compared, and the general design of systolic arrays and ahead-of-time compilation in Google TPU. Figures quoted here come from Google's Cloud TPU documentation at the time of writing; check it before you buy capacity.

Two chips, one compiler

TPU v5 ships as two different chips aimed at different jobs. v5e is the cost-efficient chip: one TensorCore per chip, 16 GB of HBM, 197 TFLOPs of peak BF16 compute and 400 GBps of bidirectional inter-chip interconnect (ICI) bandwidth, arranged in a 2D torus of up to 256 chips. Google documents it for training and single-host inference. v5p is the large-scale training chip: two TensorCores per chip, 95 GiB of HBM at 2,765 GBps, 459 TFLOPs of peak BF16, 1,200 GBps of ICI, four SparseCores per chip and a 3D torus. A single v5p slice can hold up to 6,144 chips, and Multislice extends a job to 18,432.

A TensorCore holds matrix-multiply units (MXUs), a vector unit and a scalar unit. Your code never addresses them. XLA decides which operations become MXU work, which become vector work, how to lay tensors out in HBM and where to insert collectives. The SparseCores on v5p are a separate kind of unit for the gather-heavy embedding lookups of recommendation models; dense language models do not use them.

One training step on a TPU v5 slice, and what crosses which networkYour JAX programjit(train_step)XLA compilerHLO -> fused TPU programCompilation cachekeyed by shapes + meshtrace onceHost VM 0 (one process)Input pipelinereads only its shardchipHBMchipHBMchipHBMchipHBMexecutableHost VM 1 ... NInput pipelinenext shardchipHBMchipHBMchipHBMchipHBMICI (chip to chip)Slice A = all chips joined by ICI. Gradients for tensor / FSDP sharding move here.Data-centre network (DCN)Multislice only: slice A <-> slice B, data parallelCloud Storagedataset shards, checkpointsOrbax async checkpointevery K stepsProfiler traceXProf / TensorBoard
The compiler produces one program per slice; hosts feed their own chips; ICI joins chips within a slice and DCN joins slices.

Names, hosts and provisioning

The accelerator type string is the first thing to get right, because the two chips count differently. For v5p the number is TensorCores: v5p-8 is four chips with eight TensorCores, v5p-32 is 16 chips, and v5p-128 is 64 chips. For v5e, which has one TensorCore per chip, v5litepod-16 is 16 chips. Each v5p host VM (machine type ct5p-hightpu-4t) has four chips; v5e hosts come with one, four or eight chips (ct5lp-hightpu-1t, -4t, -8t). So a v5p-128 slice is 16 host VMs, and you will run 16 copies of your program.

Each generation also needs a matching TPU software version, the image the host boots: v2-alpha-tpuv5 for v5p and v2-alpha-tpuv5-lite for v5e. Create a single slice directly, or, for anything you cannot afford to lose to a stockout, through a queued resource, which waits for capacity and creates every node at once. Note that the flag is --version on tpu-vm create but --runtime-version on queued-resources create.

# One v5p slice: 16 chips (32 TensorCores) across 4 host VMs
gcloud compute tpus tpu-vm create train-a \
    --zone=$ZONE --accelerator-type=v5p-32 --version=v2-alpha-tpuv5

# The same request, queued until capacity exists
gcloud compute tpus queued-resources create train-a-qr \
    --node-id=train-a --zone=$ZONE \
    --accelerator-type=v5p-32 --runtime-version=v2-alpha-tpuv5

# Run the same command on every host VM in the slice
gcloud compute tpus tpu-vm ssh train-a --zone=$ZONE --worker=all \
    --command="pip install -U 'jax[tpu]' && python3 train.py"

One process per host

Every host VM runs one Python process, and every process runs the same program. A process sees all devices in the slice through jax.devices(), but it can only put data on its own chips, jax.local_devices(). That leads to the first rule of TPU input pipelines: each host reads only its own share of the global batch, and the shares are stitched into one logical array that the compiled step treats as a single sharded tensor.

import jax, jax.numpy as jnp, numpy as np
from jax.sharding import Mesh, NamedSharding, PartitionSpec as P

jax.distributed.initialize()          # on Cloud TPU VMs, arguments are discovered
n_hosts, host = jax.process_count(), jax.process_index()

mesh = Mesh(np.array(jax.devices()).reshape(-1), ("data",))
batch_sharding = NamedSharding(mesh, P("data"))

GLOBAL_BATCH, SEQ = 256, 2048
LOCAL_BATCH = GLOBAL_BATCH // n_hosts  # this host's rows only

def global_batch(local_tokens: np.ndarray) -> jax.Array:
    # local_tokens: (LOCAL_BATCH, SEQ) int32, read from this host's shard files
    return jax.make_array_from_process_local_data(batch_sharding, local_tokens)

A train step that compiles once

The step function is traced once per distinct combination of input shapes, dtypes and shardings, and each trace is compiled into a TPU program. A large model can take minutes to compile, so the goal is a step that compiles exactly once. Three habits get you there. Keep shapes static: pad every sequence to a fixed length, or to a short list of bucket lengths, and drop or pad the ragged final batch. Donate the old parameters and optimizer state so XLA can update them in place instead of holding two copies in HBM. And turn on the persistent compilation cache, so a restarted job loads the compiled program instead of rebuilding it.

import functools, optax
jax.config.update("jax_compilation_cache_dir", "gs://my-bucket/jax-cache")
jax.config.update("jax_log_compiles", True)   # any compile after step 1 is a bug

opt = optax.adamw(3e-4)

def loss_fn(params, batch):
    logits = model.apply(params, batch[:, :-1])
    labels = batch[:, 1:]
    return optax.softmax_cross_entropy_with_integer_labels(logits, labels).mean()

@functools.partial(jax.jit, donate_argnums=(0, 1))
def train_step(params, opt_state, batch):
    loss, grads = jax.value_and_grad(loss_fn)(params, batch)
    updates, opt_state = opt.update(grads, opt_state, params)
    return optax.apply_updates(params, updates), opt_state, loss

for step in range(start_step, total_steps):
    batch = global_batch(next(local_iter))       # always (LOCAL_BATCH, SEQ)
    params, opt_state, loss = train_step(params, opt_state, batch)
    if step % 100 == 0:
        print(step, float(loss))                  # float() waits for the device

The float(loss) line is deliberate. JAX dispatches work asynchronously, so the Python loop normally runs ahead of the chips; converting a value to a Python number forces a wait. Do it every hundred steps, not every step, or the host stalls the device on each iteration. Sharding of the parameters themselves, for example fully sharded data parallelism, is a separate decision covered in FSDP; on TPU it is expressed by giving the parameter arrays a sharding on the mesh rather than by wrapping modules.

Worked example: feeding a v5p-128

Suppose you train a 7-billion-parameter decoder on a v5p-128 slice: 64 chips on 16 host VMs. A forward and backward pass costs about 6N FLOPs per token, so each token costs 6 x 7e9 = 4.2e10 FLOPs. Peak compute for the slice is 64 x 459e12 = 2.9e16 FLOP/s. Planning for 40 percent model FLOPs utilisation, the slice processes about 1.17e16 / 4.2e10, roughly 280,000 tokens per second.

Split across 16 hosts, each host must deliver about 17,500 tokens per second, every second, for the whole run. As raw int32 token IDs that is only 70 KB/s, trivial for the network. The cost is elsewhere: if each host is also reading compressed text, decompressing it and running a tokenizer, a few CPU cores can fall behind that rate, and the chips sit idle. The fix is to tokenize offline, store fixed-length token arrays in many shard files, give each host a disjoint set of shards and prefetch several batches ahead on a background thread.

Memory is the other half of the plan. With BF16 parameters and FP32 Adam state plus master weights, budget about 16 bytes per parameter before activations, or 112 GB for 7B parameters. Sharded across 64 chips that is under 2 GB each, which leaves most of each chip's 95 GiB for activations. The same model on a v5e slice, at 16 GB per chip, needs the state spread over many more chips or activation checkpointing; the memory-fit method is worked through in TPU v5 and v6 compared.

Multislice: crossing the data-centre network

A slice is bounded by its torus: every chip in it is joined by ICI. Multislice lets one job use several slices joined by the data-centre network (DCN), which Google's documentation describes as higher latency and lower throughput than ICI. Use it when one slice is not big enough, or when several smaller slices are easier to obtain than one large one. The rule that makes it work is to keep bandwidth-hungry parallelism inside each slice and put only data parallelism across slices, so DCN carries one gradient all-reduce per step and nothing else. In MaxText's configuration this shows up as a DCN data-parallel degree equal to the number of slices and a DCN tensor-parallel degree of 1.

Provision the slices together with a queued resource and --node-count, then build a hybrid mesh whose outer axis spans slices.

gcloud compute tpus queued-resources create big-run \
    --zone=$ZONE --accelerator-type=v5p-128 --runtime-version=v2-alpha-tpuv5 \
    --node-count=4 --node-prefix=big-run

# in train.py
from jax.experimental import mesh_utils
devices = mesh_utils.create_hybrid_device_mesh(
    mesh_shape=(1, 64),        # within a slice: 64-way FSDP over ICI
    dcn_mesh_shape=(4, 1),     # across slices: 4-way data parallel over DCN
)
mesh = Mesh(devices, ("data", "fsdp"))

Whether the DCN all-reduce hides behind computation depends on step time and gradient size; measure it in a profile rather than assuming it. Collective overlap explains the scheduling idea, which XLA applies on TPU as well.

Checkpoints, interruptions and profiling

TPU VMs do not live-migrate. A host maintenance event or a Spot preemption stops the slice, and because every host runs one program, losing one host stops the whole job. Design for restart from day one. Checkpoint with Orbax, which writes sharded arrays from every host in parallel to Cloud Storage and can do so asynchronously, so the step loop is blocked only while arrays are copied off the device. Choose the interval from the cost of lost work, and test a restore before the first long run.

import orbax.checkpoint as ocp
mngr = ocp.CheckpointManager(
    "gs://my-bucket/run-42/ckpt",
    options=ocp.CheckpointManagerOptions(save_interval_steps=2000, max_to_keep=3),
)
start_step = mngr.latest_step() or 0
if start_step:
    state = mngr.restore(start_step, args=ocp.args.StandardRestore(abstract_state))
    params, opt_state = state["params"], state["opt"]

# in the loop: a no-op except every save_interval_steps
mngr.save(step, args=ocp.args.StandardSave({"params": params, "opt": opt_state}))

The abstract_state passed to restore carries the target shardings, so a checkpoint is restored onto the current mesh. Pair the checkpoint with the data position: record the step and derive each host's shard offset from it, or a restart will replay or skip data.

For profiling, capture a few steps after warm-up with jax.profiler.trace("/tmp/jax-trace") and open the trace in TensorBoard or XProf. Read it in this order: gaps between steps mean the host or input pipeline is the bottleneck; a long tail of collective operations means communication is exposed; and a compile event after the first step means a shape changed. The general method is in GPU profiling, and it transfers almost unchanged.

Failure modes

  • Step time grows every few hundred steps. A shape is changing: a ragged last batch, a variable sequence length or a Python scalar baked into the trace. jax_log_compiles names the function; pad or bucket the input.
  • Chips idle between steps. The host cannot keep up: tokenizing on the fly, reading many small files, or a single-threaded loader. Pre-tokenize, use larger shards and prefetch.
  • Out of memory at compile time. XLA reports the program does not fit in HBM before it runs. Shard optimizer state, rematerialize activations or reduce the per-chip batch; a v5e chip has 16 GB, not 95.
  • One host hangs, all hosts hang. Collectives wait for every participant. A host that crashed or never started leaves the others blocked; check every worker's log, not just worker 0.
  • Multislice runs slower than one slice. Something other than data parallelism crossed DCN, often because the mesh axes were swapped. Print the mesh and check the outer axis is the slice axis.

Trade-offs

v5e is the cheaper chip, aimed at models that fit across many small chips, and it serves as well as trains; v5p has about six times the memory per chip, three times the ICI bandwidth and far larger slices, which matters when parameters and activations are big and communication is heavy. Choosing TPUs at all means committing to XLA: you trade hand-written kernels and dynamic shapes for a compiler that does layout, fusion and collectives for you. For dynamic shapes or unusual custom kernels that compiler becomes the bottleneck; for a large, regular transformer it removes most hand tuning.

What to do next

  1. Pick the chip by memory first: estimate bytes per parameter times parameters, divide by HBM per chip, and see which slice size fits.
  2. Create a small slice (v5litepod-8 or v5p-8) with the matching software version and run your step with jax_log_compiles on; fix every recompilation.
  3. Pre-tokenize the dataset into fixed-length shards, one disjoint set per host, and measure tokens per second per host against the worked-example target.
  4. Turn on the persistent compilation cache and Orbax checkpointing, then kill the job and prove a restart resumes at the right step and data offset.
  5. Profile five steps after warm-up and classify the time: compute, collectives, idle.
  6. Only then scale up, using a queued resource, and add Multislice with data parallelism across slices when one slice is not enough.
Key takeaway: TPU v5 runs one compiled program per slice while every host feeds its own chips. Name and provision the slice correctly, keep shapes static so the step compiles once, size the input pipeline from tokens per second per host, keep only data parallelism on the data-centre network, and checkpoint as if every run will be interrupted, because on TPU VMs it will be.