The TPU
The matrix units in a chip (MXU, vector and scalar units, on-chip memory, HBM) and the chips in a pod: an all-reduce on a 2-D torus, step by step.
Loading the animation…
Concept
A TPU pod is many chips joined by inter-chip interconnect (ICI) links into a torus: a grid whose rows and columns wrap around, so every chip has a neighbour in each direction. Training spreads a model across the chips, and after every step the chips must agree on the summed gradients: an all-reduce. The animation runs one on 16 chips in a 4 × 4 torus, each holding a vector of 16 numbers (one cell per number, darker as it sums more chips' values).
It goes dimension by dimension. Reduce-scatter along x: in each row (an x-ring), every chip passes a quarter of its vector to its +x neighbour, which adds it to its own; after steps each chip owns one quarter summed over its row. Reduce-scatter along y does the same down the columns, inside that quarter, until each chip holds one number summed over all 16 chips (outlined). All-gather along y, then x passes the finished numbers back round the same rings, copying instead of adding. The whole thing takes 12 steps; one ring through all 16 chips would take 30, though each chip sends the same 30 numbers either way. The torus wins on the number of steps, and each step's fixed cost (link latency, synchronisation) is paid fewer times: for a 16 × 16 pod, 60 steps against 510.
This is the textbook version: data moves in one direction round each ring, and one dimension at a time. Real interconnects send both ways on every link and can work on several dimensions at once, which divides the time further; the wrap-around links themselves are what make the rings possible. The TPU v4 paper reports that its 3D torus's wrap-around links double both the bisection bandwidth and the bandwidth of collective operations such as all-reduce, against a mesh without them (Jouppi et al., 2023).
Structure · static
One TPU chip, one TensorCore
Generic, after Google's public TPU documentation: four MXUs (128 × 128 systolic arrays before v6e), a vector unit and a scalar unit share on-chip memory; HBM sits beside the die; links to the neighbouring chips form the torus.
Concept
Inside one chip. Google describes each TensorCore as "one or more matrix-multiply units (MXUs), a vector unit, and a scalar unit". An MXU is a systolic array of multiply-accumulators, 128 × 128 in the TPU versions before v6e and 256 × 256 in v6e and later, multiplying bfloat16 and accumulating in FP32 (chapter 6). The vector unit does the elementwise work (activation functions, normalisation, the softmax); the scalar unit runs the control flow. Data reaches them from high-bandwidth memory (HBM) beside the die, through on-chip memory that software manages explicitly (VMEM, in the terms of JAX's Pallas TPU documentation).
The generations differ in how many of each they carry. A TPU v4 chip has two TensorCores of four MXUs each, 32 GiB of HBM2 and six ICI links, so that 4,096 chips form a 3D torus; a v5e chip has one TensorCore with four MXUs and 16 GB of HBM, and its pods are 256 chips in a 2D torus (Google Cloud TPU documentation; these change from generation to generation, so check the current pages).
The first TPU (2015, inference only) was simpler and is described in the most detail: one 256 × 256 matrix unit of 8-bit MACs, weight-stationary, fed from a 24 MiB unified buffer of activations; weights from 8 GiB of off-chip DRAM through a weight FIFO; 32-bit accumulators below the array; and an activation unit (Jouppi et al., 2017). Chapters 2, 4 and 5 are its matrix unit at small scale.
Maths
Take an torus and a vector of numbers per chip. Reduce-scatter along is steps; in step chip sends block of size to chip , which adds it. By induction, after step chip 's block holds the sum of chips, so after steps block holds all . Reduce-scatter along then splits that block ways, steps of size ; the all-gathers mirror both. Per chip that is
exactly what one ring of chips sends, and the minimum for any all-reduce (Patarasuk and Yuan, 2009). With a per-step cost and link bandwidth , a ring takes , while the torus takes : the same bandwidth term, and a latency term of instead of . The tests run every torus from 2 × 2 to 4 × 4 and check that every chip ends with the sum, that each message goes to a neighbour, and the step and message counts above.
Code
The plan for the first phase, cut from src/lib/sa/torus.ts (chip sends block to its neighbour), and the three that follow:
run("rs-x", xd - 1, (x, y, s) => [[(x + 1) % xd, y], group(mod(x - s, xd))]);
run("rs-y", yd - 1, (x, y, s) => [
[x, (y + 1) % yd],
[((x + 1) % xd) * yd + mod(y - s, yd)],
]);