Trillium is Google's sixth-generation TPU, sold on Google Cloud as TPU v6e. Google describes it as delivering 4.7 times the peak compute per chip of TPU v5e, partly by enlarging the matrix-multiply units, and as over 67 percent more energy-efficient. Those are vendor claims. What matters to someone training or serving a model is narrower: what the chip and slice actually contain, which shapes fill its matrix units, how collectives behave on its torus, and how to measure whether a program uses the hardware it is paying for.
This article covers exactly that. For a side-by-side comparison with v5e and v5p, ridge-point arithmetic and fitting models into HBM, read TPU v5 and v6 compared. For how systolic arrays and the XLA compiler work in general, read Google TPU. Figures here come from Google's v6e and TPU architecture documentation as of October 2026; check them again before you plan capacity.
Trillium in one picture
What Google documents
Google's v6e page lists the following per chip. Note what is absent: Google does not publish a clock speed, and the documented MXU count and peak figure do not reconcile with an obvious clock, so this article derives nothing from them. Treat 918 TFLOPs as a ceiling to measure against.
| Per chip | TPU v6e (Trillium) |
|---|---|
| Peak BF16 compute | 918 TFLOPs |
| Peak Int8 compute | 1,836 TOPs |
| HBM capacity and bandwidth | 32 GB at 1,638 GBps |
| Inter-chip interconnect (ICI) | 800 GBps, 4 ports |
| Compute units | 1 TensorCore with 2 MXUs, a vector unit and a scalar unit |
| MXU size | 256 x 256 multiply-accumulators (128 x 128 on earlier generations) |
| Other | SparseCore for embedding lookups |
| Pod | 256 chips in a 2D torus; 102.4 TB/s all-reduce, 3.2 TB/s bisection |
For comparison, v5e is listed at 197 TFLOPs and 16 GB of HBM, which is where the 4.7 times figure comes from. The change that most affects how you write software is the MXU: it doubled in each dimension, so each unit holds four times as many multiply-accumulators.
How a matrix multiplication meets a 256 x 256 array
An MXU is a systolic array: a grid of multiply-accumulate cells through which operands flow in lockstep, each cell multiplying, adding to a running sum and passing values to its neighbour. Weights for one tile are loaded into the grid, activations stream through, and partial sums accumulate in FP32 from bfloat16 inputs. A matrix multiplication of an M by K activation with a K by N weight is cut by XLA into tiles that fit the array, and the K and N dimensions of each weight tile map onto the grid's 256 by 256 cells.
When a dimension is not a multiple of the tile size, the last tile is padded with zeros, and those cells do no useful work. On a 128-wide array, a dimension of 128 fills it; on a 256-wide array, the same dimension fills half. The helper below computes the useful fraction of a padded dimension.
import math
def mxu_fill(dim: int, tile: int = 256) -> float:
"""Fraction of a padded MXU dimension doing useful work."""
return dim / (math.ceil(dim / tile) * tile)
for name, d in [("hidden 4096", 4096), ("hidden 5120", 5120),
("head_dim 128", 128), ("vocab 50257", 50257),
("LoRA rank 16", 16)]:
print(f"{name:16s} 256-tile fill {mxu_fill(d):.2f}"
f" 128-tile fill {mxu_fill(d, 128):.2f}")Hidden sizes of 4,096 and 5,120 are multiples of 256 and fill the array completely. A GPT-2-style vocabulary of 50,257 pads to 50,432, a negligible loss, though padding the vocabulary yourself to a multiple of 256 is free. The losses that matter are small dimensions: an attention head dimension of 128 in the score computation, or a LoRA rank of 16. Whether XLA recovers some of that by packing several small operations together depends on the operation and compiler version, which is why the next step is to measure.
Measuring instead of estimating
Never take a utilisation estimate on trust. A microbenchmark of a plain matrix multiplication at shapes that are and are not multiples of 256 shows on your own slice how much padding costs, and the profiler trace shows what XLA actually emitted.
import time
import jax
import jax.numpy as jnp
PEAK = 918e12 # documented v6e BF16 peak, per chip
def achieved_tflops(m, k, n, iters=50):
a = jnp.ones((m, k), jnp.bfloat16)
b = jnp.ones((k, n), jnp.bfloat16)
f = jax.jit(lambda x, y: x @ y)
f(a, b).block_until_ready() # compile once, outside timing
t0 = time.perf_counter()
for _ in range(iters):
out = f(a, b)
out.block_until_ready()
secs = (time.perf_counter() - t0) / iters
return 2 * m * k * n / secs
# 4096 and 3840 are multiples of 256; 3968 only of 128; 4000 of neither.
for k in (4096, 3840, 3968, 4000):
flops = achieved_tflops(8192, k, 8192)
print(k, f"{flops / 1e12:.0f} TFLOPs", f"{flops / PEAK:.0%} of peak")
with jax.profiler.trace("/tmp/v6e-trace"): # open in TensorBoard / XProf
achieved_tflops(8192, 4096, 8192, iters=5)Run it on one chip of a v6e-1 or v6e-8 VM. The gap between the 256-multiple rows and the others is the padding and tiling cost in practice. Then profile one real training step: the trace shows each fused operation's time, and a step where large matrix multiplications account for a small share of time is bound by something else, usually memory-bound elementwise work, the input pipeline or collectives. Model FLOPs utilisation, achieved model FLOPs divided by 918 TFLOPs times the number of chips, is the single number to track per run.
Workloads that underfill the array
Three workloads routinely waste a 256-wide array, and each has a known response. Decoding. Here the problem is not padding: each step processes one token per sequence, so every weight loaded from HBM is used for only a few tokens, and the step is memory-bound. Batch more sequences per chip, using continuous batching, before adding chips; the ridge-point arithmetic is in the comparison article linked above. LoRA. Adapter matrices with ranks of 8 to 64 waste most of the array, but they are a small share of FLOPs during training; for serving a single adapter, merge it into the base weights so it costs nothing. Mixture of experts. Tokens per expert per step can be small and uneven; larger groups and capacity factors keep per-expert matrices wide enough, a trade-off discussed in expert parallelism.
Hosts, VMs and the input pipeline
Chips are attached to hosts in groups of eight, and you get them as VMs of one, four or eight chips. Google's table gives the VM shapes:
| VM | vCPUs | RAM | NUMA nodes |
|---|---|---|---|
| 1 chip | 44 | 176 GB | 1 |
| 4 chips | 180 | 720 GB | 1 |
| 8 chips | 360 | 1,440 GB | 2 |
Slices range from 1x1 to 16x16. A 4x4 slice is 16 chips on 4 VMs; a full 16x16 pod slice is 256 chips on 64 VMs. Each VM runs one copy of your program, and every copy must call jax.distributed.initialize() before touching devices. Each host loads only its own share of the data, so the input pipeline runs on 45 vCPUs per chip; tokenising text on the fly or decoding images can starve the chips, which the profile shows as idle gaps between steps. On the eight-chip VM, keep data-loading threads on the NUMA node local to their chips where your framework allows it.
Collectives on the 16 x 16 torus
Each chip has four ICI ports, one to each neighbour in the 2D torus, with wraparound at the edges. Google gives two pod-level figures that bound collective performance: 102.4 TB/s of all-reduce bandwidth and 3.2 TB/s of bisection bandwidth. Dividing by 256 chips gives about 400 GB/s per chip for all-reduce, while the bisection figure divided across the 128 chips on one side of the cut is only 25 GB/s per chip. The ms figures below are estimates from these aggregates, not measurements.
All-reduce, used for data-parallel gradients and for the gather and reduce-scatter steps of FSDP, moves roughly twice the buffer size per chip. For 16 GB of bf16 gradients from an 8-billion-parameter model, that is about 32 GB at roughly 400 GB/s, so on the order of 80 ms, or nearer 40 ms if Google's figure already counts that factor of two.
All-to-all, used by expert parallelism, sends data from every chip to every other chip, so about half of it crosses any cut through the middle of the torus. Across a full 256-chip slice, its throughput is bounded by the bisection, and per chip that is more than ten times lower than the all-reduce rate. The rule that follows: keep all-to-all and tensor parallelism inside small groups of neighbouring chips, and put the axis that only needs all-reduce across the whole slice. Lay out the device mesh so its axes match the torus.
import jax
from jax.experimental import mesh_utils
from jax.sharding import Mesh, NamedSharding, PartitionSpec as P
jax.distributed.initialize() # on every VM of a multi-host slice
# 256 chips as 16 x 16, matching the physical torus. "data" spans rows,
# "model" spans the 16 chips of each row, so model traffic stays on short paths.
devices = mesh_utils.create_device_mesh((16, 16))
mesh = Mesh(devices, ("data", "model"))
weights = NamedSharding(mesh, P(None, "model"))
batch = NamedSharding(mesh, P("data", None))The general theory of these collectives, written for GPUs but applicable here, is in NCCL collectives.
Provisioning and launching a slice
Multi-host slices are usually obtained through queued resources: the request waits in a queue until the capacity exists, and then the whole slice is created at once. The runtime version for v6e is v2-alpha-tpuv6e. Because every VM must run the same program, launch with --worker=all, or through a job system such as GKE that does the equivalent.
# Request a 16-chip slice through the queue (preferred for multi-host slices).
gcloud alpha compute tpus queued-resources create my-qr \
--node-id=my-v6e --project=${PROJECT_ID} --zone=${ZONE} \
--accelerator-type=v6e-16 --runtime-version=v2-alpha-tpuv6e
# Once ACTIVE, run the same command on every VM of the slice.
gcloud compute tpus tpu-vm ssh my-v6e --zone=${ZONE} --worker=all \
--command="python3 train.py"A multi-host job fails when any VM fails, so checkpoint to Cloud Storage at a fixed interval and make restart automatic before running anything long.
Worked example: an 8B model on v6e-256
Take an 8-billion-parameter dense transformer trained on a full v6e-256 slice with a global batch of about 4.2 million tokens, so 16,384 tokens per chip per step. Training costs about 6 FLOPs per parameter per token, so each chip does 6 x 8e9 x 16,384, about 7.9e14 FLOPs per step. At an assumed 45 percent model FLOPs utilisation, the chip delivers about 413 TFLOPs, so a step takes about 1.9 seconds. That is about 2.2 million tokens per second for the slice, or roughly 5 days for a trillion tokens.
Now the communication. With FSDP across all 256 chips, each step all-gathers the 16 GB of bf16 weights for the forward pass, again for the backward pass, and reduce-scatters 16 GB of gradients: on the order of 48 GB per chip, which at roughly 400 GB/s is about 0.12 seconds, around 6 percent of the step, and XLA can overlap most of it with compute. The conclusion is that this configuration is compute-bound, so the work that pays is raising utilisation: fixed, bucketed sequence lengths to avoid recompilation, dimensions that are multiples of 256, and a profile showing large fused matrix multiplications dominating each step. Turn the same model into a mixture of experts with all-to-all across the full slice, and the bisection estimate above says communication, not compute, will set the pace.
Failure modes
- Shapes that half-fill the array. Dimensions of 128 that were ideal on v5e use half of each v6e MXU dimension. Benchmark before assuming a v5e configuration carries over.
- Recompilation. Every new input shape compiles a new program. Bucket sequence lengths and batch sizes, and watch the compile count.
- Starved hosts. 45 vCPUs per chip is generous but finite; heavy preprocessing on the fly shows up as gaps in the profile. Preprocess offline.
- All-to-all across the whole slice. Bisection bandwidth caps it; keep expert groups local.
- Out of HBM. 32 GB per chip holds weights and optimizer state easily under FSDP; activations are usually what overflows. Use rematerialisation before adding chips.
- Whole-job failure. One lost VM stops every VM. Checkpoint often and automate restart.
- Cross-slice traffic. Multislice traffic leaves ICI for the data-centre network; put only data parallelism across slices.
Trade-offs
Trillium's large MXUs give very high peak compute per chip, and the price is sensitivity to shape: small and ragged matrices waste more of each unit than they did on 128-wide arrays. Its 2D torus is cheap and fast for neighbour traffic and all-reduce, but its bisection is modest, which favours dense models with data and FSDP parallelism over communication-heavy all-to-all across a whole pod. SparseCore helps embedding-heavy ranking models and does nothing for dense language models. And the XLA programming model trades flexibility for performance: fixed shapes and whole-program compilation are what let the compiler fill the arrays.
What to do next
- Read the current v6e specifications and topology table, and note the slice shape and VM type you need.
- Run the matrix-multiplication benchmark on one chip and record achieved TFLOPs for 256-multiple and other shapes.
- Pad model dimensions, vocabulary and batch sizes to multiples of 256 where it is cheap, and bucket sequence lengths.
- Profile one real training or serving step and compute model FLOPs utilisation against 918 TFLOPs per chip.
- Build the device mesh to match the 16 x 16 torus, with all-to-all and tensor parallelism inside small groups.
- Estimate collective time per step from the pod aggregates, then confirm it in the profile.
- Provision through queued resources, launch on all workers, and set up checkpointing and automatic restart before long runs.