systolic-arrays-explained
← /learn · 07

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 X−1X - 1 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.

ICI −yICI +y−x+xTensorCoreMXU 1systolic arrayMXU 2systolic arrayMXU 3systolic arrayMXU 4systolic arrayVector unitelementwise: activations, softmaxScalarcontrolOn-chip memory(VMEM)HBMbeside the die, in the package

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 X×YX \times Y torus and a vector of DD numbers per chip. Reduce-scatter along xx is X−1X - 1 steps; in step ss chip xx sends block (x−s) mod X(x - s) \bmod X of size D/XD/X to chip x+1x + 1, which adds it. By induction, after step ss chip xx's block (x−s−1) mod X(x - s - 1) \bmod X holds the sum of s+2s + 2 chips, so after X−1X - 1 steps block (x+1) mod X(x + 1) \bmod X holds all XX. Reduce-scatter along yy then splits that block YY ways, Y−1Y - 1 steps of size D/(XY)D/(XY); the all-gathers mirror both. Per chip that is

sent=2(X−1)DX+2(Y−1)DXY=2D XY−1XY,\text{sent} = 2(X - 1)\frac{D}{X} + 2(Y - 1)\frac{D}{XY} = 2D\,\frac{XY - 1}{XY},

exactly what one ring of XYXY chips sends, and the minimum for any all-reduce (Patarasuk and Yuan, 2009). With a per-step cost α\alpha and link bandwidth β\beta, a ring takes 2(XY−1)(α+D/(XYβ))2(XY - 1)(\alpha + D/(XY\beta)), while the torus takes 2(X−1)(α+D/(Xβ))+2(Y−1)(α+D/(XYβ))2(X - 1)(\alpha + D/(X\beta)) + 2(Y - 1)(\alpha + D/(XY\beta)): the same bandwidth term, and a latency term of 2(X+Y−2)α2(X + Y - 2)\alpha instead of 2(XY−1)α2(XY - 1)\alpha. 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 xx sends block (x−s) mod X(x - s) \bmod X to its +x+x 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)],
]);