Google's Cloud TPU v5 and v6 generations give you three different machines, not one. TPU v5e is a cost-efficient chip for training and serving small and medium models. TPU v5p is the large-scale training chip, with the most memory per chip and pods of thousands of chips. TPU v6e, called Trillium, is the sixth generation, with much more compute per chip than v5e in a similar 256-chip pod. Choosing between them and using them well is a question of arithmetic: how many FLOPs, how many bytes, and how far data has to travel.
This article sets out the published specifications, explains how software sees the hardware, and shows how to turn the numbers into decisions with roofline and memory-fit calculations, JAX sharding code and topology-aware parallelism. The general ideas behind TPUs, systolic arrays and ahead-of-time compilation are covered in Google TPU. Figures here come from Google's Cloud TPU documentation; check the current pages before committing to a purchase, because offerings change.
The three chips side by side
| Per chip | v5e | v5p | v6e (Trillium) |
|---|---|---|---|
| Peak BF16 compute | 197 TFLOPs | 459 TFLOPs | 918 TFLOPs |
| Peak Int8 compute | 393 TOPs | Not compared here | 1,836 TOPs |
| HBM capacity | 16 GB | 95 GiB | 32 GB |
| HBM bandwidth | 800 GiBps | 2,765 GBps | 1,638 GBps |
| ICI bandwidth (bidirectional) | 400 GBps | 1,200 GBps | 800 GBps |
| TensorCores | 1 (4 MXUs) | 2 | 1 (2 MXUs) |
| Topology | 2D torus | 3D torus | 2D torus |
| Chips per pod | 256 | 8,960 (largest job: 6,144) | 256 |
Google describes Trillium as delivering 4.7 times the peak compute per chip of v5e, achieved by enlarging the matrix-multiply units and raising the clock, and as being over 67% more energy-efficient than v5e. Those are vendor claims; the table's figures show the compute ratio directly (918 against 197). Trillium also includes Google's third-generation SparseCore, and v5p has four SparseCores per chip. SparseCores accelerate the large embedding lookups in ranking and recommendation models; dense language-model training does not use them.
Google has since announced a further generation, Ironwood. It is outside the scope of this article; check its documentation separately if you are choosing hardware today.
What is on the chip, and how software sees it
A TensorCore holds matrix-multiply units (MXUs), which do the dense matrix work, a vector unit for elementwise operations such as activations and normalization, and a scalar unit that sequences everything. High-bandwidth memory (HBM) sits next to the chip, and inter-chip interconnect (ICI) links connect each chip directly to its neighbours, without going through a network switch.
You do not program these units directly. You write JAX, or PyTorch through PyTorch/XLA, and the XLA compiler turns the whole step function into a fused program, choosing memory layouts, padding shapes to suit the MXUs and inserting the collective operations that sharding requires. Each host runs one process that sees the chips attached to it as devices, and in a multi-host slice every host runs the same program.
Because XLA compiles for exact shapes, every new input shape triggers a new compilation. Pad batches and sequences to a small set of bucket sizes. The MXU dimensions are not published on Google's specification pages; dimensions that are multiples of 128 are a safe default on all three chips, and larger multiples are worth benchmarking on v6e, whose MXUs are widely reported to be larger.
Accelerator types and slices
You rent TPUs as slices: a set of chips connected by ICI in a particular shape. The accelerator-type name counts different things on different generations, which is a common source of mistakes.
- v5p:
v5p-Ncounts TensorCores, two per chip, sov5p-8is 4 chips. Slice shapes are 3D, from2x2x1(4 chips) up to16x16x24(6,144 chips), built from 4x4x4 cubes. - v5e:
v5litepod-Ncounts chips;v5litepod-16is a 4x4 slice. Shapes run from 1x1 to 16x16. Google documents serving support only on 1x1, 2x2 and 2x4 slices. - v6e:
v6e-Ncounts chips; shapes run from 1x1 to 16x16 (256 chips on 64 VMs). Single VMs hold 1, 4 or 8 chips.
# Create a single 8-chip v6e TPU VM (values from Google's v6e training guide).
gcloud alpha compute tpus tpu-vm create my-tpu \
--version=v2-alpha-tpuv6e \
--accelerator-type=v6e-8 \
--zone=${ZONE} \
--project=${PROJECT_ID}Training beyond one slice uses Multislice, which joins several slices over the data-centre network. Traffic inside a slice uses ICI; traffic between slices does not, and is far slower, which shapes how you assign parallelism, as discussed below. Large slices are often obtained through queued resources or through GKE node pools rather than by creating VMs directly.
Roofline arithmetic: which chip is your workload bound on?
Every operation either waits on compute or waits on memory. The ridge point of a chip is its peak FLOPs divided by its memory bandwidth: the number of FLOPs an operation must do per byte it reads to keep the MXUs busy. From the table:
| Chip | Peak BF16 / HBM bandwidth | Ridge point |
|---|---|---|
| v5e | 197 TFLOPs / ~859 GB/s (800 GiBps) | ~230 FLOPs per byte |
| v5p | 459 TFLOPs / 2,765 GB/s | ~166 FLOPs per byte |
| v6e | 918 TFLOPs / 1,638 GB/s | ~560 FLOPs per byte |
Now take the core of a transformer layer, multiplying activations for B tokens by a d by d weight matrix in bf16. It does 2Bd² FLOPs and reads 2d² bytes of weights, so its intensity is about B FLOPs per byte when the weight read dominates. To be compute-bound, B must exceed the ridge point: roughly 166 tokens per step on v5p, 230 on v5e and 560 on v6e.
This has direct consequences. During training, batches of thousands of tokens per chip are normal, so all three chips run compute-bound and v6e's extra FLOPs translate into speed. During autoregressive decoding, each step processes one token per sequence, so B is the number of concurrent sequences. With 64 sequences in flight, all three chips are memory-bound, and the chip with the most HBM bandwidth per dollar wins. v6e only pays off for serving when you batch heavily. The same analysis is developed in the roofline model.
Worked example: fitting models into HBM
Training an 8B-parameter model with Adam. A common mixed-precision layout keeps bf16 weights (2 bytes per parameter), bf16 gradients (2 bytes) and fp32 master weights plus two Adam moments (12 bytes): 16 bytes per parameter, so 128 GB before activations. That does not fit on any single chip. Sharding parameters, gradients and optimizer state across chips, the approach known as FSDP, divides it: across 8 v6e chips (256 GB in total) each chip holds 16 GB of state and keeps 16 GB for activations, which is workable with activation checkpointing. On v5e, with 16 GB per chip, the same model needs at least 16 chips and more for comfortable headroom. A v5p-8, which is 4 chips with 95 GiB each, holds the state with plenty of room to spare.
Serving a 70B-parameter model in bf16. The weights take 140 GB. A v6e-8 slice has 256 GB, leaving about 100 GB for the KV cache and runtime buffers. A 2x4 v5e slice has only 128 GB, which is too small without int8 weights; quantizing to int8 halves the weights to 70 GB and uses the v5e's int8 throughput. These estimates ignore framework overheads, so confirm with a real run and a memory profile.
Sharding code in JAX
JAX expresses parallelism by placing devices in a named mesh and annotating how arrays are split across its axes. XLA inserts the collectives. The example builds a two-axis mesh for FSDP plus tensor parallelism and jits a training step.
import jax
import numpy as np
from jax.sharding import Mesh, NamedSharding, PartitionSpec as P
devices = np.array(jax.devices()) # one entry per chip on v5e and v6e
mesh = Mesh(devices.reshape(-1, 4), ("fsdp", "tp")) # e.g. 2 x 4 on v6e-8
def shard(spec):
return NamedSharding(mesh, spec)
# Weights: rows split over fsdp, columns over tp. Batch split over fsdp.
w_sharding = shard(P("fsdp", "tp"))
x_sharding = shard(P("fsdp", None))
@jax.jit
def train_step(params, opt_state, batch):
loss, grads = jax.value_and_grad(loss_fn)(params, batch)
updates, opt_state = optimizer.update(grads, opt_state, params)
params = jax.tree_util.tree_map(lambda p, u: p + u, params, updates)
return params, opt_state, loss
params = jax.device_put(params, w_sharding) # in practice, a per-leaf sharding tree
batch = jax.device_put(batch, x_sharding)On a multi-host slice the same script runs on every host after jax.distributed.initialize(), and each host loads only its share of the data. Libraries such as MaxText package these patterns for language models. The parameter-sharding idea itself is explained in FSDP.
Topology-aware parallelism
Different kinds of parallelism produce different traffic. Tensor parallelism exchanges activations inside every layer and needs the fastest links. FSDP gathers weights and reduces gradients once per layer per step, and can overlap that with compute. Data parallelism across replicas exchanges gradients once per step. The rule is to put the most communication-hungry axis on the best links.
Inside a slice, everything rides ICI, and a torus gives every chip direct neighbours in each dimension, with wraparound links in a full torus. v5p's 3D torus offers more paths and three times v5e's ICI bandwidth, which is why it scales to thousands of chips for very large models. v5e and v6e are capped at 256 chips per slice; beyond that, use Multislice and put only data parallelism across slices, since that traffic happens once per step and tolerates the slower data-centre network.
Choosing a generation
| Workload | Usual first choice | Reason |
|---|---|---|
| Serving a small or quantized model at low cost | v5e | Lowest-cost chip; serving supported on small slices |
| Serving larger models with heavy batching | v6e | 32 GB per chip and high compute once batches are large |
| Training and fine-tuning up to a few hundred chips | v6e | Highest compute per chip in a 256-chip pod |
| Training very large models on thousands of chips | v5p | 95 GiB per chip, 3D torus, jobs up to 6,144 chips |
| Embedding-heavy recommendation models | v5p or v6e | Both include SparseCores |
Price per chip-hour and regional capacity change often, and they matter as much as the specifications. Benchmark your own step time on two candidates and compare cost per useful token, not peak FLOPs.
Failure modes
- Recompilation storms. Variable sequence lengths or batch sizes trigger compiles that take seconds to minutes each. Bucket shapes and watch compile counts.
- Starved input pipeline. A slow data loader on the host leaves chips idle. Profile step time against host CPU time.
- Out of HBM. Activations, not weights, usually push training over the limit. Use activation checkpointing and smaller micro-batches before buying more chips.
- Wrong slice arithmetic. Ordering
v5p-8and expecting 8 chips gives you half the memory you planned for. - Lost capacity. Preemptible or spot capacity can disappear, and multi-host jobs fail when any host does. Checkpoint frequently to Cloud Storage and make restart automatic.
- Cross-slice bottleneck. Putting tensor parallelism across Multislice boundaries sends per-layer traffic over the data-centre network and collapses throughput.
What to do next
- Compute your model's memory footprint per chip for training or serving, and shortlist the generations where it fits with headroom.
- Compute tokens per step per chip and compare it with each chip's ridge point to see whether you will be compute-bound or memory-bound.
- Port the training or serving step to JAX or PyTorch/XLA with fixed, bucketed shapes, and confirm the compile count stays flat.
- Run the same step on a small slice of two generations, for example v5e and v6e, and compare cost per token.
- Lay out the device mesh so tensor parallelism stays inside a slice, and use Multislice only for data parallelism.
- Set up checkpointing and automatic restart before scaling up, and check current regional availability for the slice shape you need.