A Cloud TPU pod is a supercomputer that you rent in pieces. The pieces are slices: rectangular blocks of chips wired to each other by a dedicated inter-chip interconnect (ICI) in a torus, with no switch between neighbours. That wiring is why TPUs scale the way they do, and it is also why choosing a slice shape, and mapping your model's parallelism onto it, has a large effect on step time.

This article treats the pod as a network. It defines the vocabulary, compares the current pod generations, shows how torus wraparound changes collective cost, builds a small cost model, and maps a JAX device mesh onto ICI and the data-centre network (DCN). Chip internals live in TPU v6 Trillium and TPU v4 and v5; provisioning lives in GKE with TPU.

Pods, slices, cubes and hosts

Google's documentation uses precise terms; using them precisely avoids sizing mistakes.

  • Pod: a contiguous set of TPU chips grouped over a specialised network.
  • Slice: chips inside one pod connected by ICI. You allocate slices, not pods. The topology is the slice's shape, such as 4x4x4 (3D) or 16x16 (2D).
  • Cube: a 4x4x4 block of 64 chips, the building unit of 3D-torus generations. Inside a cube, ICI uses copper; cubes are joined through optical circuit switches.
  • Host: a VM attached to a few chips. Each host runs one process of your program; a multi-host slice runs the same program on every host.
  • DCN: the ordinary data-centre network. Slices talk to each other over it; Multislice is the feature that spans one job over several slices.

One naming trap: v5p accelerator types count TensorCores, two per chip, so v5p-128 is a 4x4x4 slice of 64 chips. v6e types count chips, so v6e-256 is 256 chips. Read the topology, not the suffix, when budgeting.

Three pod generations

Three current generations, from Google's Cloud TPU documentation as of October 2026:

v5pv6e (Trillium)TPU7x (Ironwood)
Chips per pod8,9602569,216
ICI topology3D torus2D torus3D torus
HBM per chip95 GiB32 GB192 GiB
ICI per chip, bidirectional1,200 GBps800 GBps1,200 GBps
Chips per VM44 per VM in multi-host slices4

The shape of the network matters as much as the numbers. A v6e pod is a 16 x 16 grid: each chip has four neighbours, and the largest slice is the whole pod, so larger jobs must use Multislice over DCN. v5p and Ironwood are 3D, with six neighbours per chip and slices of thousands of chips; Google documents 6,144 chips (96 cubes) as the largest single-slice v5p job.

The torus and wraparound

In a torus, each axis is a ring: the last chip in a row links back to the first. Rings matter because ring collectives use every link at once. Google documents a detail that changes performance: on v5p, slices of a full cube or larger have full 3D torus connectivity, but slices smaller than a cube are 3D connected without wraparound links. Their axes are lines, not rings.

A 4x4 slice of a torus: rings per axis, with and without wraparoundSmaller than a full cube: lines, no wraparoundFull axis: rings, wraparound links (dashed)ring on a line uses each link once per directiontwo rings, one per direction: about 2x per axisMesh axis to physical axis: put the heaviest collective on the axis with wraparounddata across slices rides DCN; tensor and FSDP shards stay on ICI
Without wraparound, a ring collective must be embedded in a line, using each link once per direction. With wraparound, two independent rings run in opposite directions, roughly doubling per-axis collective bandwidth.

Some shapes also come twisted: the wraparound links connect to an offset position, which shortens the longest paths. Google reports about 70 percent higher bisection bandwidth for twisted 4x4x8 slices than for the plain torus, which helps all-to-all traffic such as embedding lookups and mixture-of-experts dispatch more than it helps all-reduce.

A cost model for collectives on a torus

A useful model of a collective on a torus treats each mesh axis separately. An all-reduce over a multi-dimensional slice is done axis by axis: reduce-scatter along x, then y, then z, each step on a smaller shard, then all-gather back. Each axis costs about 2(k-1)/k times the data on that step, divided by the bandwidth of the rings along it.

def axis_allreduce_s(nbytes, k, link_GBps, wrap):
    # ring all-reduce along one axis of length k; two rings if wraparound exists
    if k == 1:
        return 0.0
    rings = 2 if wrap else 1
    return 2 * (k - 1) / k * nbytes / (rings * link_GBps * 1e9)

def torus_allreduce_s(nbytes, dims, link_GBps, wrap):
    # reduce-scatter then all-gather per axis; data shrinks by k after each axis
    t, shard = 0.0, nbytes
    for k in dims:
        t += axis_allreduce_s(shard, k, link_GBps, wrap)
        shard /= k
    return t

# Illustration only: assume ~100 GB/s per link per direction.
grads = 16e9   # 8B parameters, bf16 gradients
print(torus_allreduce_s(grads, (4, 4, 4), 100, wrap=True))   # full cube
print(torus_allreduce_s(grads, (2, 4, 4), 100, wrap=False))  # sub-cube

The per-link figure is an assumption, not a published number: dividing the documented 1,200 GBps per chip across six links and two directions gives 100 GB/s, but achieved bandwidth depends on the compiler's collective implementation. The model's job is relative comparisons, wraparound versus none and one shape versus another; measure absolute numbers with a profile.

Mapping a JAX mesh onto ICI and DCN

In JAX you never address ICI links. You build a logical device mesh with named axes, annotate arrays with how they are split across those axes, and the XLA compiler inserts collectives; see XLA compilation. Performance depends on which physical axis each logical axis lands on.

import jax
from jax.sharding import NamedSharding, PartitionSpec as P
from jax.experimental import mesh_utils

jax.distributed.initialize()          # one process per host
# Single slice, 64 chips: data x fsdp x tensor
mesh = jax.make_mesh((4, 4, 4), ("data", "fsdp", "tensor"))

# Two slices over DCN: data parallel across slices, everything else on ICI
devices = mesh_utils.create_hybrid_device_mesh(
    mesh_shape=(1, 8, 8),             # per slice (ICI)
    dcn_mesh_shape=(2, 1, 1))         # across slices (DCN)
mesh2 = jax.sharding.Mesh(devices, ("data", "fsdp", "tensor"))

w = NamedSharding(mesh2, P("fsdp", "tensor"))   # weight matrix sharding
x = NamedSharding(mesh2, P(("data", "fsdp"), None))  # batch over data and fsdp

Two rules follow from the cost model. Put the axis with the most traffic per step, usually tensor parallelism, on ICI links that form rings. And put only the lightest traffic on DCN: pure data parallelism, which communicates once per step and overlaps with backward, is the conventional choice across slices.

Multislice and the DCN budget

Multislice exists because slices have ceilings: 256 chips on v6e, and on v5p whatever contiguous block of cubes the scheduler can find. Joining slices over DCN raises the ceiling but changes the network from a torus of fast neighbour links to a shared datacentre fabric with much less bandwidth per chip. Google quotes 25.6 Tbps of DCN bandwidth per v6e pod; spread over 256 chips that is about 100 Gbps, or 12.5 GB/s, per chip, against 800 GBps of ICI per chip. That ratio, more than sixty to one, is the number to remember.

It follows that cross-slice traffic must be small relative to the step and must overlap with compute. A practical budget:

def dcn_budget(grad_bytes_per_slice, chips_per_slice, dcn_GBps_per_chip, step_s):
    # data-parallel all-reduce across slices: each chip moves its shard over DCN
    shard = grad_bytes_per_slice / chips_per_slice
    t = 2 * shard / (dcn_GBps_per_chip * 1e9)   # ~2x data for a 2-slice ring
    return t, t / step_s

# 8B model, bf16 grads, 256-chip v6e slices, 12.5 GB/s DCN per chip, 1.5 s step
t, frac = dcn_budget(16e9, 256, 12.5, 1.5)
print(f"{t*1000:.0f} ms of DCN per step, {frac:.0%} of the step")   # ~10 ms, ~1%

A fully sharded gradient reduction across slices is cheap because each chip moves only its shard. The same arithmetic applied to per-layer FSDP all-gathers over DCN is ruinous, which is the quantitative reason the mesh rule above exists.

Worked example: one large slice or two cubes

Suppose you train an 8B-parameter dense model on v5p and can get either one 4x4x8 slice (128 chips, v5p-256) or two 4x4x4 slices joined by DCN. Each chip has 95 GiB of HBM; with FSDP sharding, weights, gradients and Adam state (about 16 bytes per parameter, 128 GB in total) spread across 64 chips come to 2 GB per chip, leaving room for activations at sequence length 8,192.

On the single slice, use a mesh of (data=2, fsdp=64) or (fsdp=128) entirely on ICI. On two slices, use (data=2 over DCN, fsdp=64 over ICI). The FSDP all-gathers and reduce-scatters, the heavy per-layer traffic, stay inside a cube with wraparound either way; only one gradient all-reduce of the 64-way shards crosses DCN, 16 GB / 64 = 250 MB per chip per step, and it overlaps with backward.

Then measure: profile a step in each configuration and compare the time in collective ops. If DCN time is exposed, increase the batch per step or accumulate gradients so the cross-slice reduction amortises over more compute.

Reading a step profile

The cost models give expectations; a profile gives facts. Capture a few steps with jax.profiler.trace("/tmp/trace") after warm-up, open the trace in TensorBoard or Perfetto, and sum the time in collective ops (all-gather, reduce-scatter, all-reduce, all-to-all) separately from matrix work. Three readings are common. Collectives that run alongside compute and do not lengthen the step are fine. Collectives in a long serial block at the end of the step point to the last gradient reduction being exposed, often across DCN. And all-to-all dominating the trace on an MoE model is a cue to try a twisted topology or a different expert layout. Record the per-step collective time for each topology you use, so you can spot a degraded slice by comparison.

Failure modes

  • Slice as failure unit. A slice runs as one program; a chip or host failure stops the whole slice. Checkpoint on a cadence that matches the slice's interruption rate, and plan restarts as normal operations.
  • Sub-cube slices without wraparound give lower collective bandwidth than their chip count suggests. Benchmark before choosing a small 3D shape.
  • Wrong axis on DCN. Putting tensor or FSDP axes across slices sends per-layer traffic over the data-centre network and can make the step several times slower. Check the mesh with jax.devices() and each device's slice_index.
  • Recompilation on shape change. A different topology means a new compiled program. Cache compilations and keep shapes fixed between restarts.
  • Host-side input starvation. Each host feeds only its own chips. If one host's input pipeline lags, every chip waits at the next collective.
  • Optical faults. Google documents ICI resiliency that routes around optical-switch and optical-link faults; it keeps jobs running but may lower bandwidth, so a slowly degrading step time warrants a support case.

Trade-offs

A pod's switchless torus gives very high neighbour bandwidth at low cost and power, and a compiler that schedules collectives ahead of time. The cost is rigidity: slices are fixed shapes, bandwidth depends on shape and wraparound, all-to-all traffic is harder than on a switched fat tree, and you buy into XLA. One big slice keeps all traffic on ICI but is harder to obtain and fails as one unit; several smaller slices with Multislice are easier to schedule and replace but put data-parallel traffic on DCN. For a broader comparison with GPU clusters, see Google TPU.

What to do next

  1. Translate every accelerator type into chips and topology before budgeting; remember v5p suffixes count TensorCores.
  2. Prefer full-cube 3D slices, or confirm with a benchmark that a sub-cube slice without wraparound is fast enough.
  3. Lay out your JAX mesh so tensor and FSDP axes stay on ICI and only data parallelism crosses DCN.
  4. Run the cost model for each candidate shape, then profile one step on each.
  5. Consider twisted topologies for all-to-all heavy models such as MoE and embeddings.
  6. Checkpoint with a cadence matched to slice interruptions, and test restore on a different slice.
  7. Track per-step time in collectives over time; a slow climb can mean degraded links.
Key takeaway: A TPU pod is a switchless torus that you rent as slices. Collective speed depends on slice shape and on whether its axes have wraparound links, which on v5p needs a full cube. Map heavy parallelism onto ICI rings, keep only data parallelism on DCN, count chips rather than suffixes, and measure each candidate shape before committing.