Tiling big GEMMs
A matrix bigger than the array, cut into weight tiles and run back to back: double-buffered weights hide each tile's load, and the dataflow decides how many words the buffer must deliver.
Loading the animation…
Concept
A real layer is much bigger than the array: a 1024 × 1024 weight matrix on a 128 × 128 array is 64 tiles. Chapter 4 ran tiles one after another, each paying its own weight load, fill and drain. This chapter overlaps them, as the TPU does.
The trick is a second, shadow weight register in every PE. While tile 1's activations stream through the array, tile 2's weights shift down the columns into the shadow registers, one row per cycle. When the first activation of tile 2 reaches a PE, the PE swaps: the shadow weight becomes the one it multiplies by. That first activation enters on the cycle after tile 1's last one, so the swap ripples across the array as a diagonal, exactly behind the end of tile 1's wavefront. Tile 1's last partial sums drain out of the bottom while tile 2's first ones start at the top.
In the animation is 6 × 6 and is 6 × 6, on a 3 × 3 array: 4 weight tiles (two of by two of ). Back to back they take 31 cycles (77% of the array's MAC slots busy); in sequence they took 52 (46%). Switch the weight registers to single to see the difference. With a short stream ( below the array's height) the load cannot hide: a tile cannot start before its weights are in, so the gap between tiles becomes cycles. At , the decode-like case, back to back still beats in sequence (17 cycles against 32), but most of the array idles.
The model does more than draw this. Every cycle it checks that each PE swaps in the weight of the right tile and that every partial sum meets its next product in step; a test moves one tile's load or stream a single cycle earlier and the model stops with an error.
Concept
The weights have to come from somewhere. Set HBM words per cycle to a number and each tile's 9 weights must first be fetched from off-chip memory into an on-chip weight buffer (purple bars). With one buffer, the next fetch can only start once the previous tile has been shifted into the array, and the fetches show up in the total: 64 cycles at one word per cycle. With two buffers (ping-pong: fetch into one while the array reads the other) the fetch overlaps the stream: 49 cycles, and 36 at two words per cycle. To hide a fetch completely the memory must deliver a tile's weights while the previous tile streams for cycles: 1.5 words per cycle here.
This is the first TPU's design. Its weights came from off-chip DRAM through a weight FIFO, four tiles deep, and its matrix unit held two tiles of weights, so that the 256 cycles needed to shift a tile in were hidden behind the previous tile's computation (Jouppi et al., 2017, section 2). Its activations lived in a 24 MiB unified buffer on the chip.
Loading the animation…
Concept
Inside the chip the question is the same one chapter 1 asked of off-chip memory: how many words must the on-chip buffer deliver per MAC? Tiling decides it. When a GEMM is cut into tiles, each operand is read again for every tile of the dimension it does not span: in weight-stationary, is streamed once per column tile of , and a tile that covers only part of writes partial sums that must be read back and added (the sky bars). Output-stationary keeps each result in its PE until it is complete, so it writes once, but reads once per row tile of .
For a layer with 1,024 and rows, an 8 × 8 array needs 0.376 words per MAC in weight-stationary and 0.251 in output-stationary; a 256 × 256 array needs 0.0127 and 0.008789. Every doubling of the array roughly halves the traffic per MAC: a bigger array reuses each word more. Switch to (decode-like: few tokens) and weight-stationary bottoms out at 0.1357 words per MAC on the big array, because every weight is read for only 8 rows. Small batches starve big arrays of reuse, whatever the dataflow.
Maths
Tile of a weight-stationary run covers rows of and columns , and streams all rows of . Write for the cycle its first activation enters the array and for the cycle its weights start loading. Column loads in cycles (the load is skewed like the activations), and PE swaps at , when the tile's first activation arrives. The shadow register of PE is first overwritten by the next load in cycle , so the next load may start once this tile has started, ; the weights are all in place cycles after it starts. Two tiles may not feed the same row in the same cycle, so . Hence
with the end of the tile's fetch ( when the weights are already on chip), and the run ends when the last tile's last result leaves: . With the weights on chip and , every tile after the first costs exactly cycles: the array never stops. The tests check this closed form, and the tiles-in-sequence total of chapter 4 when the shadow registers are switched off, against the cycle-by-cycle model.
For the buffer traffic, with an array and dimensions that divide evenly, weight-stationary reads times ( words), once (), writes partial or final results and reads back partial sums. Dividing by the MACs gives the equations above; the widget uses the exact counts (tiledWords), which the Python tests check against the cycle-accurate tiled simulation.
Code
The schedule, cut from src/lib/sa/stream.ts (a port of stream_schedule in the Python reference):
let load: number;
if (j === 0) load = ready;
else if (shadow) load = Math.max(out[j - 1]!.stream, ready);
else load = Math.max(out[j - 1]!.end, ready);
let stream = j === 0 ? load + kj : Math.max(out[j - 1]!.stream + m, load + kj);
and the swap inside a PE, from the cycle-by-cycle model:
let s = old[0];
if (hIn !== null && (s === null || s[3] !== hIn[3])) {
s = old[3];
/* c8 ignore next */
if (s === null || s[3] !== hIn[3] || s[1] !== hIn[2])
throw new Error("weight not loaded in time");
}