systolic-arrays-explained
← /learn · 05

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 AA is 6 × 6 and BB is 6 × 6, on a 3 × 3 array: 4 weight tiles (two of KK by two of NN). 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 (MM below the array's height) the load cannot hide: a tile cannot start before its weights are in, so the gap between tiles becomes kk cycles. At M=1M = 1, 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 knkn weights while the previous tile streams for max⁡(M,k)\max(M, k) 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, AA is streamed once per column tile of BB, and a tile that covers only part of KK 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 CC once, but reads BB once per row tile of AA.

For a layer with K=N=K = N = 1,024 and M=512M = 512 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 M=8M = 8 (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 jj of a weight-stationary run covers rows k0…k0+kj−1k_0 \ldots k_0 + k_j - 1 of BB and columns n0…n0+nj−1n_0 \ldots n_0 + n_j - 1, and streams all MM rows of AA. Write SjS_j for the cycle its first activation enters the array and LjL_j for the cycle its weights start loading. Column qq loads in cycles Lj+q…Lj+q+kj−1L_j + q \ldots L_j + q + k_j - 1 (the load is skewed like the activations), and PE(k,q)(k, q) swaps at Sj+k+qS_j + k + q, when the tile's first activation arrives. The shadow register of PE(k,q)(k, q) is first overwritten by the next load in cycle Lj+1+q+kL_{j+1} + q + k, so the next load may start once this tile has started, Lj+1≥SjL_{j+1} \ge S_j; the weights are all in place kj+1k_{j+1} cycles after it starts. Two tiles may not feed the same row in the same cycle, so Sj+1≥Sj+MS_{j+1} \ge S_j + M. Hence

Sj+1=max⁡(Sj+M, Lj+1+kj+1),Lj+1=max⁡(Sj, Fj+1),S_{j+1} = \max\bigl(S_j + M,\ L_{j+1} + k_{j+1}\bigr), \qquad L_{j+1} = \max(S_j,\ F_{j+1}),

with Fj+1F_{j+1} the end of the tile's fetch (00 when the weights are already on chip), and the run ends when the last tile's last result leaves: T=SJ+M+kJ+nJ−2T = S_J + M + k_J + n_J - 2. With the weights on chip and M≥kM \ge k, every tile after the first costs exactly MM 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 R×CR \times C array and dimensions that divide evenly, weight-stationary reads AA N/CN/C times (MKN/CMKN/C words), BB once (KNKN), writes MN⋅K/RMN \cdot K/R partial or final results and reads back MN(K/R−1)MN(K/R - 1) partial sums. Dividing by the MKNMKN 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");
}