Retrieved article excerpt
Open article · Retrieved 2026-09-17T05:21:38.575269+00:00
[← Andrew Fan](https://im-afan.github.io/)
# Designing my own TPU for transformer inference and making it go brr
## Introduction
This was a summer project I made to learn the basics of ML performance optimization and hardware-software co-design. The overall goal of this project is to run a small transformer on my Cmod A7 FPGA board. Then, I want to derive some of the common results you’ll find in LLM scaling (limited to a single device for now), and optimize my kernels to see how much performance I can squeeze out of my hardware design. I’m assuming you have basic knowledge on how transformer inference works (prefill, decode), but not the hardware side of things. All the code can be found in the project’s [github repo](https://github.com/im-afan/transformer-tpu-int4).
Much of the analysis I do (including the next section), along with the TPU architecture itself, is heavily inspired by the [Scaling Book](https://jax-ml.github.io/scaling-book/). If you find this article interesting, you might enjoy reading the scaling book!
## Background
### What determines how fast our model is?
Before we get into hardware, we need to ask a seemingly obvious question: what makes a model/algorithm faster? In ML, this problem can be analyzed in 2 domains: compute and communications. Compute is how many raw operations we can do in some amount of time. For example, FLOP/s (floating point operations / sec) is a common metric for the compute capabilities of GPUs or other ML accelerators. So the total compute time of an algorithm can be calculated using \(T\_{\text{compute}} = \frac{\text{Total FLOPs}}{\text{Accelerator FLOPs/s}}\).
On the other hand, we have communication. In single-device inference, this usually refers to the comms between the memory (HBM, DDR) and the accelerator’s cache. In distributed inference, the comms between devices must also be considered. Similarly, if we know the memory bandwidth of our HBM or DDR, we can calculate our comms time as \(T\_{\text{comms}} = \frac{\text{Communication Bytes}}{\text{Memory Bytes/sec}}\)
In most hardware, we assume that comms and compute run at the same time, so optimally, they are completely overlapped. So, our lower bound on the runtime of an algorithm is \(\max(T\_{\text{comms}}, T\_{\text{compute}})\). Generally speaking, our goal in designing both hardware and kernels is to analyze the bottleneck between these two and optimize it. For example, it wouldn’t make sense to try to speed up a memory-bound algorithm by trying to use more arithmetic units!
### Matmuls
Matmuls have a special property regarding compute and comms that makes them the backbone of ML. First, let us consider an \(N\times N\times N\) matmul, in int8 (our design uses int4, but int8 keeps the byte counts clean): \(A[N, N] \cdot B[N, N] = C[N, N]\). In the best case, we have to load \(2N^2\) bytes to our compute, then write back \(N^2\) bytes to memory, which is a total of \(3N^2\). For compute, we have to perform \(N^3\) int8 multiplications. So the total arithmetic intensity of the algorithm is on the order of N, meaning it is more compute-demanding the larger our matmul gets. This is partly why matmuls are so essential to ML workloads: it is very easy to scale up our models by just throwing more compute at it to achieve larger matmuls.
## Model Architecture / Goal
The overall goal is to run a transformer model on our design, while not making it completely fixed to a single architecture. We want to be able to write code to change the model architecture/dimensions, tweak our matmuls to be more efficient, or run unit tests to benchmark single layers or matmuls. This means that having fixed control flow set in hardware is not acceptable.
We will train our model and run inference on predicting the next token in an addition sequence with a maximum of 31 digits per input (64 token prefill and 64 token decode overall), which is a pretty simple but nontrivial task.
For benchmarking, we use a standard Transformer architecture, with embedding dim \(d=128\), ffn dim \(d\_{ff}=512\), 4 QKV heads, and 4 layers. However, we greatly simplify some operations to allow for ease of implementation. First, we replace gelu with relu activation. Softmax in attention is also replaced with a ReLU, which makes training much harder to converge but is fine for our task. LayerNorm is replaced with a hardtanh normalization. Finally, we train with no bias in the ffn. While these are pretty major simplifications, the overall architecture stays the same, and the original goal of analyzing transformer inference performance can still be achieved.
## Hardware
### Overview
In a transformer, we have 2 types of operations: matmuls, and elementwise ops like tensor addition and ReLU. The purpose of an accelerator is to load tensors from memory, do those operations, and write back the results. To make the most of our FPGA resources and memory, we choose to quantize activations and weights to int4.
The memory hierarchy of our TPU is simple. We are using a Cmod A7 board, which includes an Artix-7 FPGA chip along with an external asynchronous SRAM chip (8 bit read, 10 ns access time). We use the SRAM chip to model our accelerator’s external memory (HBM/DDR in a real accelerator), and the Artix-7’s BRAM to act as an on-chip cache (scratchpad memory).
There are 3 main units: MXU, VPU, and DMA. The MXU (matrix multiply unit) handles the matrix multiplications, VPU (vector processing unit) handles vector operations such as activations and addition. Both units read from the scratchpad memory and write their results back to scratchpad. The DMA (direct memory access) handles transfers between scratchpad and external memory. Each unit is controlled by a softcore PicoRV32 processor, for which we can write C firmware to dispatch instructions to each unit through AXI-based MMIO.
TPU architecture
*Abstract layout of a real TPU TensorCore, whose shape ours follows. Source: [How To Scale Your Model](https://jax-ml.github.io/scaling-book/tpus/)*
There are a few differences in our design. Vmem is our scratchpad, and HBM is our external SRAM chip. The scalar unit is just the PicoRV32 core. We also have a DMA unit, which is what actually moves data between the SRAM and the scratchpad.
### Memory
Our scratchpad memory is synthesized as simple dual-port BRAM, with 2 independent ports: read and write. Since it can be synthesized to basically an arbitrary bus width for our usage (1 x 32K bits to 72 x 512 bits), we don’t need the blocks to be traditional banks indexed by `address % bus_width`. Instead, we can use a single block to represent a contiguous region of our scratchpad memory. This basically allows us to have as many ports accessing scratchpad as we want, with the limitation that 2 ports don’t access the same region; this will be very useful in the design later.
Simple dual-port BRAM
*A 7-series block RAM in simple dual-port mode: one read port and one write port over the same 36 Kb array. Source: [Xilinx UG473, 7 Series FPGAs Memory Resources](https://docs.amd.com/v/u/en-US/ug473_7Series_Memory_Resources)*
Our DMA is very simple. Running at the Cmod A7’s standard 12 MHz clock, the 10 ns external RAM access time is basically instant; it arrives at the next clock. The DMA unit takes in a base address for scratchpad and external RAM, a row stride, and 2d matrix dimensions, and simply drives the ports of the external RAM and scratchpad to transfer data between them. However, there is one wrinkle: since our external SRAM is asynchronous and has a 10 ns access time, the WE (write-enable) signal needs to come only after the address and data signals have arrived and stabilized. This means that for each byte written, we need to turn WE on and back off again, at the negative edge of our clock. So our final bandwidth is 1 byte / clock for external -> scratchpad, and 0.5 byte / clock for scratchpad -> external.
### MXU & VPU
For matmuls, we use an output-stationary 8x8 systolic array; each PE keeps its partial sum, while moving its input values to the next PE to its right and below it. This allows us to perform an arbitrary 8xNx8 (\(A[8,N] \cdot B[N,8]\)) matmul, as long as A and B fit in scratchpad. By default, all tensors are stored row-major in both external memory and scratchpad. Since our design has no fast way to transpose a matrix, we instead use a trick in the MXU to perform transposed matmuls such as \(QK^T\) in attention. By default, A and B are fed into the systolic array like this:
MXU dataflow
*2x2 version of the MXU; the real one is 8x8. Each PE stores one output element \(c\_{ij}\). Left, no transpose: row \(i\) of A (blue) is loaded into that row’s register as one contiguous chunk and shifts out one element per clock. A whole B row arrives at once, and registers skew the incoming data. Right, transposed: B is stored \([N, K]\) instead, so a column’s elements are contiguous and B is fed the same way as A. This makes it so that transposed matrices never need to be re-arranged in memory; the MXU just reads them differently.*
As a result, every clock, the systolic array reads \(32\) bits (8 int4 values) for matrix A and \(32\) bits for matrix B. This is where the scratchpad architecture comes in handy: as long as A and B are in different memory regions, they can be read at the same time, without interfering with ongoing DMA operations. With all \(64\) PEs accumulating one product each clock, our MXU does \(64\) multiply-accumulates per clock, which is the compute number we will use for the rest of this article.
The VPU is pretty simple. Its inputs are the base address of a vector and the length of the operation. It repeatedly loads chunks of the vector from scratchpad, operates on them (add, relu, etc.), and writes them back, until the operation is completed.
### Design Choices & Other Notes
We use PicoRV32 so that we can easily write firmware for different architectures. Not only does this allow for unit testing beyond just inference and different architectures, it allows us to easily experiment with optimizations later on without having to change the dataflow in hardware.
CPU issue overhead was also a concern when designing the architecture. When the CPU dispatches instructions, they enter a command queue for each unit, which are then executed asynchronously from the CPU execution order. For synchronization, the CPU can also poll each instruction queue’s state. Furthermore, to minimize the effect of execution latency, the TPU operations are intentionally complex, allowing for instructions to span across large address ranges without needing more CPU dispatch calls.
## Firmware & Kernels
We now need to program the PicoRV32 core to perform inference. First, we define our programming scheme:
- Since our external SRAM (512 KB) is much larger than our scratchpad (64 KB), we want to only use the scratchpad to handle the individual operands of a primitive. Every operation should read and write back to external memory.
- We will define many primitives such as matmul, tensor addition, ReLU, and then compose them together using RISCV’s loop & branching capabilities to implement inference.
We start by implementing the basic instruction dispatch to the TPU. After that, we implement primitives, namely matmul and elementwise functions.
### Optimizing A Matmul
We are lucky enough to have a relatively big scratchpad (64 KB) compared to our external SRAM (512 KB), which is a 1:8 ratio. With d=128, d\_ff=512, and T=64 during prefill, our largest matmul, in the FFN, requires about \((128 \* 512 + 64 \* 128 + 64 \* 512) / 2 = 53,248\) bytes (52 KB), which can fit entirely in our scratchpad! This means that for the matmuls in our benchmark, we can basically always achieve the theoretical \(NM + MK + NK\) byte loads, since we don’t need to load the same chunk of a matrix twice when tiling.
We implement matmul as a tiled matmul over \(8\times 8\) tiles in the output matrix. Here’s the pseudocode for th