systolic-arrays-explained
← /learn · 09

From graph to silicon

An ONNX model lowered onto the array: im2col turns a convolution into a GEMM, the GEMMs become weight tiles, and the cycle-accurate count meets a simulator's estimate.

Loading the animation…

Concept

A model arrives at an accelerator as a graph of operators, here an ONNX file of four nodes: a convolution, a ReLU, a flatten and a fully connected (Gemm) layer. Before any of it can run on a systolic array, a compiler must lower each node to what the hardware does: matrix multiplies for the array, elementwise work for the vector unit, and nothing at all for operations that only reshape.

The convolution becomes a GEMM. The trick is im2col: for every output pixel, copy the patch of input it sees (every channel, the whole R×SR \times S window) into one row of a matrix AA. Flatten each filter into one column of BB. Then the convolution is ABA B: row pp of AA times column ff of BB is output pixel pp of filter ff. For this model's 36 output pixels, 3 input channels and 3 × 3 filters that is a 36 × 27 by 27 × 8 GEMM. The animation builds AA one row at a time; the copying costs memory (each input pixel appears in up to 9 rows) but turns the convolution into the one thing the array does well. The site runs this GEMM on its cycle-accurate array and the tests check it against a direct convolution, number for number.

The GEMMs become tiles. On an 8 × 8 array the convolution's K=K = 27 needs 4 weight tiles, run back to back as in chapter 5: 161 cycles, 75% of the array busy. The fully connected layer is the opposite case. Its input is a single row (one image), so every one of its 72 weight tiles is loaded to multiply one row of AA: 585 cycles and 7.7% busy, waiting on weight loads. This is the batch-of-one problem of LLM decoding in miniature, and why accelerators batch requests. ReLU's 288 elements go to the vector unit; Flatten moves no data.

Concept

Three ways to count the cycles. The bars under the animation give three numbers for each array layer. Back to back is this site's cycle-accurate model with double-buffered weights (chapter 5). In sequence runs the same tiles one after another (chapter 4): 222 cycles for the convolution. Approximate is the formula in the author's Torch_Sim_Frontend, a fast event-driven model of a whole accelerator (DMA, buffers, array, vector unit), which charges an output-stationary array kk cycles per block of outputs plus one fill and drain: 151 cycles for the convolution, 592 for the fully connected layer.

The approximate formula and the cycle-accurate model agree to within a few per cent on the fully connected layer (592 against 586 for output-stationary tiles in sequence) and differ by more on the convolution, where the formula overlaps every tile's fill and drain while the in-sequence count pays for each. That is the trade a performance model makes: one line of arithmetic per tile instead of a register-level simulation, so it can run whole networks quickly, at the price of an error that this kind of cross-check measures. scripts/check_simfront.py runs Torch_Sim_Frontend's own ONNX front end on the same file (at commit de3acd4): its GEMM shapes and its cycle formula agree with this page's.

Maths

For a convolution with CinC_{\text{in}} input channels of H×WH \times W, CoutC_{\text{out}} filters of R×SR \times S, stride 1 and no padding, the output is Cout×Hout×WoutC_{\text{out}} \times H_{\text{out}} \times W_{\text{out}} with Hout=H−R+1H_{\text{out}} = H - R + 1 and Wout=W−S+1W_{\text{out}} = W - S + 1. im2col builds

Ap, (c,r,s)=xc,  ip+r,  jp+s,B(c,r,s), f=wf,c,r,s,p=ipWout+jp,A_{p,\,(c, r, s)} = x_{c,\; i_p + r,\; j_p + s}, \qquad B_{(c, r, s),\, f} = w_{f, c, r, s}, \qquad p = i_p W_{\text{out}} + j_p,

so that (AB)p,f=∑c,r,sxc,ip+r,jp+s wf,c,r,s=of,ip,jp(AB)_{p,f} = \sum_{c,r,s} x_{c, i_p + r, j_p + s}\, w_{f,c,r,s} = o_{f, i_p, j_p}. The GEMM has M=HoutWoutM = H_{\text{out}} W_{\text{out}} (times the batch), K=CinRSK = C_{\text{in}} R S and N=CoutN = C_{\text{out}}: here 36×27×836 \times 27 \times 8, 7,776 MACs. Torch_Sim_Frontend's estimate for a GEMM on an R×CR \times C output-stationary array is

T≈batch⋅⌈M/R⌉⋅⌈N/C⌉⋅K+R+C,T \approx \text{batch} \cdot \lceil M / R \rceil \cdot \lceil N / C \rceil \cdot K + R + C,

and the cycle-accurate counts are those of chapters 4 and 5.

Code

im2col, cut from src/lib/sa/lower.ts (a port of reference/lower.py, which reads the ONNX file):

for (let i = 0; i < h - r + 1; i++)
  for (let j = 0; j < w - s + 1; j++) {
    const row: number[] = [];
    for (let ci = 0; ci < c; ci++)
      for (let di = 0; di < r; di++)
        for (let dj = 0; dj < s; dj++) row.push(x[ci]![i + di]![j + dj]!);
    rows.push(row);
  }