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.

Advertisement

The three chips side by side

Per chipv5ev5pv6e (Trillium)
Peak BF16 compute197 TFLOPs459 TFLOPs918 TFLOPs
Peak Int8 compute393 TOPsNot compared here1,836 TOPs
HBM capacity16 GB95 GiB32 GB
HBM bandwidth800 GiBps2,765 GBps1,638 GBps
ICI bandwidth (bidirectional)400 GBps1,200 GBps800 GBps
TensorCores1 (4 MXUs)21 (2 MXUs)
Topology2D torus3D torus2D torus
Chips per pod2568,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.

How a JAX program reaches TPU hardware, and where each generation differsJAX / PyTorch-XLAjit, sharding annotationsXLA compilerfusion, layout, collectivesTPU runtime (PJRT)one process per hostOne chipTensorCoreMXUs + vector + scalarSparseCoreembedding lookupsHBMv5e 16 GB / v6e 32 GB / v5p 95 GiBICI linksto neighbouring chipsv5e and v6e: 2D torusup to 256 chips per podv5p: 3D torus8,960 chips per podAcross slices (Multislice), traffic leaves ICI and uses the data-centre network, which is much slower.
Frameworks hand a whole program to XLA, which compiles it for the chip. Each chip has a TensorCore, SparseCores, HBM and ICI links. v5e and v6e pods are 2D tori of up to 256 chips; v5p pods are 3D tori.

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.

Advertisement

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-N counts TensorCores, two per chip, so v5p-8 is 4 chips. Slice shapes are 3D, from 2x2x1 (4 chips) up to 16x16x24 (6,144 chips), built from 4x4x4 cubes.
  • v5e: v5litepod-N counts chips; v5litepod-16 is a 4x4 slice. Shapes run from 1x1 to 16x16. Google documents serving support only on 1x1, 2x2 and 2x4 slices.
  • v6e: v6e-N counts 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:

ChipPeak BF16 / HBM bandwidthRidge point
v5e197 TFLOPs / ~859 GB/s (800 GiBps)~230 FLOPs per byte
v5p459 TFLOPs / 2,765 GB/s~166 FLOPs per byte
v6e918 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

WorkloadUsual first choiceReason
Serving a small or quantized model at low costv5eLowest-cost chip; serving supported on small slices
Serving larger models with heavy batchingv6e32 GB per chip and high compute once batches are large
Training and fine-tuning up to a few hundred chipsv6eHighest compute per chip in a 256-chip pod
Training very large models on thousands of chipsv5p95 GiB per chip, 3D torus, jobs up to 6,144 chips
Embedding-heavy recommendation modelsv5p or v6eBoth 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-8 and 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

  1. Compute your model's memory footprint per chip for training or serving, and shortlist the generations where it fits with headroom.
  2. 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.
  3. Port the training or serving step to JAX or PyTorch/XLA with fixed, bucketed shapes, and confirm the compile count stays flat.
  4. Run the same step on a small slice of two generations, for example v5e and v6e, and compare cost per token.
  5. Lay out the device mesh so tensor parallelism stays inside a slice, and use Multislice only for data parallelism.
  6. Set up checkpointing and automatic restart before scaling up, and check current regional availability for the slice shape you need.
Key takeaway: TPU v5e, v5p and v6e are three different tools. v5e is the low-cost chip for small and medium models, v5p has the most memory and the largest 3D-torus pods for very large training runs, and v6e (Trillium) has the most compute per chip in a 256-chip pod. Choose with arithmetic: memory fit per chip, tokens per step against each chip's ridge point, and which parallelism axis needs ICI. Then write sharded JAX or PyTorch/XLA code with fixed shapes and benchmark cost per token on the candidates.