systolic-arrays-explained
← /learn · 02

Weight-stationary, cycle by cycle

The TPU's dataflow: weights sit still, skewed activations march right, partial sums flow down and drain out of the bottom. Every register, every cycle.

Loading the animation…

Concept

This is the whole machine, one clock cycle per step, with the value in every register. The array is K×NK \times N PEs: one PE for every weight wkn=bknw_{kn} = b_{kn} of BB.

  1. Load (the first KK cycles). The weights enter at the top, one row of BB per cycle, and shift down until PE(k,n)(k, n) holds wknw_{kn} (the vermillion "w" in each PE). From then on they stay put: the dataflow is named after what is stationary.
  2. Stream. Row mm of AA enters from the left, but not all at once: amka_{mk} enters row kk of the array at compute cycle m+km + k. This skew, the staircase of blue values waiting at the left edge, is what makes everything meet at the right time. Each activation steps one PE to the right per cycle, so it is used by all NN PEs of its row.
  3. Accumulate downwards. Each PE multiplies the activation passing through by its weight and adds the product to the partial sum that arrived from the PE above (sky "p"), then passes the new sum down. PE(k,n)(k, n) handles row mm at cycle m+k+nm + k + n; the PE above did so one cycle earlier, so its partial sum is exactly the one needed.
  4. Drain. After KK PEs the sum is complete, and cmnc_{mn} leaves the bottom of column nn (purple), at compute cycle m+(K−1)+nm + (K - 1) + n.

With the default problem (AA is 5 × 4, BB is 4 × 4, on 16 PEs), the run takes 15 cycles: 4 to load and 11 to stream, compute and drain. The first result, c00=c_{00} = 25, leaves in cycle 8 and the last in cycle 15. The array does 80 useful MACs, 33.3% of what its PEs could have done in that time; chapter 4 is about where the rest goes.

Change MM, KK or NN: the model re-runs and the animation restarts. More rows of AA use the same loaded weights again, which is why the TPU holds weights: in inference a weight matrix is reused by every token in a batch.

Maths

Write τ\tau for the compute cycle (after the KK load cycles). The left edge of row kk receives

ink(τ)={aτ−k, k0≤τ−k<Mbubbleotherwise.\text{in}_k(\tau) = \begin{cases} a_{\tau - k,\,k} & 0 \le \tau - k < M \\ \text{bubble} & \text{otherwise.} \end{cases}

Every register is a flip-flop, so a value latched by PE(k,n)(k, n) at the end of cycle τ\tau reaches PE(k,n+1)(k, n+1) (activations) or PE(k+1,n)(k+1, n) (partial sums) in cycle τ+1\tau + 1. By induction on nn, amka_{mk} is in PE(k,n)(k, n) in cycle m+k+nm + k + n; by induction on kk, the partial sum for (m,n)(m, n) holding kk terms arrives at PE(k,n)(k, n) in the same cycle m+k+nm + k + n. So PE(k,n)(k, n) computes

pmn(k+1)=pmn(k)+amk wkn,pmn(0)=0,p^{(k+1)}_{mn} = p^{(k)}_{mn} + a_{mk}\,w_{kn}, \qquad p^{(0)}_{mn} = 0,

and pmn(K)=∑kamkbkn=cmnp^{(K)}_{mn} = \sum_k a_{mk} b_{kn} = c_{mn} leaves the bottom at τ=m+K−1+n\tau = m + K - 1 + n. The last result, cM−1,N−1c_{M-1,N-1}, leaves at τ=M+K+N−3\tau = M + K + N - 3, so a run takes

TWS=K⏟load+(M+K+N−2) cycles.T_{\text{WS}} = \underbrace{K}_{\text{load}} + (M + K + N - 2)\ \text{cycles}.

The tests check every one of these claims against the frames: each product amkwkna_{mk} w_{kn} happens in PE(k,n)(k, n) at cycle m+k+nm + k + n, each result leaves when complete, and CC equals numpy's matmul.

Code

The heart of the model, cut from src/lib/sa/model.ts: one PE in one compute cycle. prev is the previous cycle's registers, so everything moves one PE per cycle by construction (ws covers weight- and input-stationary, which are the same machine with AA and BB swapped):

const wVal = cell[0]![0];
cell[4] = [hIn[0], wVal];
cnt.macs += 1;
cell[2] = [pVal + hIn[0] * wVal, mi, ni, terms + 1];
cnt.regWrites += 1;
if (kk === k - 1) {
  out.push([mi, ni, cell[2][0]]);
  cnt.writes += 1;
}

The Python reference has the same loop; the TypeScript port reproduces every register of every PE in every cycle exactly, for every problem size the sliders allow.