Modern AI Engineering

Lesson 17.3 · 26 min

Google TPUs: Purpose-Built Hardware for Neural Networks

What if, instead of building a chip that can do anything, we built one that does almost nothing except multiply matrices — but does that thousands of times more efficiently?

In short: A TPU (Tensor Processing Unit) is Google's custom chip for neural networks. Its heart is a systolic array: a grid of simple multiply-add cells through which data pulses in lockstep, so each number fetched from memory is reused many times without being re-read. That makes matrix multiplication very fast and power-efficient, at the cost of flexibility: TPUs shine on large, regular, compiled workloads and struggle with irregular or dynamic ones.

What is a TPU?

A TPU (Tensor Processing Unit) is an ASIC — an application-specific integrated circuit, a chip designed for one job — that Google builds to run neural networks. A 'tensor' is just a multi-dimensional array of numbers (a vector is a 1-D tensor, a matrix a 2-D tensor), and neural networks are mostly operations on tensors.

Google started using its first TPU inside its data centres in 2015 and revealed it in 2016. Since then it has released many generations: later ones train as well as run models, use high-bandwidth memory, and are connected by the thousands into pods. TPUs power Google products such as Search, Photos and Translate, were used to train Google's Gemini models, and are rented to other companies through Google Cloud.

Think of it like a dedicated pasta machine A home kitchen (a CPU) can cook anything. A restaurant kitchen with many cooks (a GPU) can cook many dishes at once. A pasta factory machine (a TPU) only makes pasta — but it turns flour into noodles faster, cheaper and with less energy than any kitchen. If our menu is 90% pasta, the factory machine is a great investment. If the menu changes every day, it is not.

Why Google built the TPU

Around 2013, Google saw deep learning spreading through its products — speech recognition especially. Google has described an internal projection from that time: if people used voice search for just a few minutes a day, running the neural networks on the CPUs of the day would have required roughly doubling its data-centre capacity. Buying that many general-purpose servers was not realistic.

So Google built a chip that did the core neural-network operation far more efficiently. The first TPU was designed, built and deployed in about 15 months. In Google's 2017 paper on it, the TPU v1 ran Google's production inference workloads roughly 15–30× faster than the contemporary CPU and GPU it was compared with, and delivered roughly 30–80× better performance per watt.

TPU generations (high level)

  1. TPU v1: Inference only. 8-bit integer maths on a 256×256 systolic array (65,536 multiply-accumulate cells) at 700 MHz, with 28 MiB of on-chip memory.
  2. TPU v2: Adds training: floating-point maths with the bfloat16 format (which Google introduced), high-bandwidth memory, and chips linked into pods.
  3. v3 and v4: More compute per chip, liquid cooling (v3), and much larger pods (v4 links thousands of chips with reconfigurable optical switches).
  4. v5e, v5p, Trillium, Ironwood: Variants tuned for cost-efficiency or peak scale, and newer generations with larger matrix units and memory; Google positioned Ironwood (announced 2025) especially for inference.

Numbers change every generation Exact figures — matrix-unit size, memory, bandwidth, chips per pod — differ by generation and are best checked in Google Cloud's current documentation. The design ideas in this lesson are what stay the same.

A quick refresher: CPU and GPU

The one operation that matters most

Look inside any modern neural network and one operation dominates: matrix multiplication, built from multiply-accumulate (MAC) steps: acc = acc + a × b. Fully connected layers, the attention projections and feed-forward blocks of a Transformer, and convolutions (which can be rewritten as matrix multiplies) all reduce to enormous numbers of MACs.

Here is the key observation behind the TPU. In a matrix multiply Y = X · W, every input value is used many times: each entry of X is multiplied by a whole row of W, and each weight is used for every row of X. A general processor tends to fetch values from memory (or at least from registers) again for each use. But moving data costs far more energy than multiplying it. Estimates for chips of the 2010s put an off-chip DRAM access at hundreds of times the energy of an 8- or 16-bit multiply. So the winning design is one that fetches each number once and reuses it as many times as possible.

Pause and think: For Y = X · W with X of size 2×3 and W of size 3×3, how many multiply-accumulate operations are needed, and how many times is each element of X used?

2 × 3 × 3 = 18 MACs. Each element X[m, k] is multiplied by every weight in row k of W, so it is used 3 times (once per output column).

The big idea: the systolic array

A systolic array is a grid of identical, very simple processing elements (PEs). Each PE can do one multiply-accumulate per clock tick and talks only to its immediate neighbours. Data enters at the edges and moves one PE per tick, in rhythm — like blood pumped by a heart, which is where the name 'systolic' comes from. The idea dates back to H. T. Kung and Charles Leiserson's work in the late 1970s; the TPU applied it at a huge scale.

Think of it like a bucket brigade Instead of each firefighter running to the well, people stand in a line and pass buckets hand to hand. Each bucket is filled once at the well and travels the whole line. In a systolic array, each number is read from memory once and handed from PE to PE, being used at every stop.

The TPU v1 used a weight-stationary design. Before the computation, a block of weights is loaded so that each PE holds one weight, W[k, n]. Then:

  • Activations (the input values X) flow in from the left edge and move one PE to the right each tick.
  • Partial sums flow downward one PE each tick. At each PE: new partial sum = partial sum from above + (activation passing through × the PE's stored weight).
  • Finished results fall out of the bottom edge: each column produces one output column of Y.

How data flows through a TPU

Let us simulate a tiny 3×3 weight-stationary array computing Y = X · W, where X has two input rows. The inputs must be skewed (staggered in time): row k of the array receives its value one tick later than row k − 1, so that each activation meets the right partial sum at the right PE.

systolic_array.py

import numpy as np
# Compute Y = X @ W on a tiny 3x3 weight-stationary systolic array.
X = np.array([[1, 2, 3],
[4, 5, 6]])            # 2 inputs (M=2) of size K=3
W = np.array([[1, 0, 2],
[0, 1, 1],
[1, 1, 0]])            # K=3 x N=3 weights, one per PE
M, K = X.shape; N = W.shape[1]
act = np.zeros((K, N), int)     # activation register in each PE (moves right)
psum = np.zeros((K, N), int)    # partial-sum register in each PE (moves down)
Y = np.zeros((M, N), int)
for t in range(M + K + N - 2):
# 1) every activation hops one PE to the right; new ones enter on the left,
#    skewed so row k receives X[m, k] at cycle m + k
act[:, 1:] = act[:, :-1].copy()
for k in range(K):
m = t - k
act[k, 0] = X[m, k] if 0 <= m < M else 0
# 2) every partial sum hops one PE down and picks up act * weight
above = np.vstack([np.zeros((1, N), int), psum[:-1]])
psum = above + act * W
# 3) finished sums fall out of the bottom row: column n carries row m=t-(K-1)-n
for n in range(N):
m = t - (K - 1) - n
if 0 <= m < M:
Y[m, n] = psum[K - 1, n]
print(f"cycle {t}: activations entering left = {act[:, 0].tolist()}")
print("systolic result:\n", Y)
print("numpy X @ W:\n", X @ W)
print("match:", np.array_equal(Y, X @ W), f"| cycles: {M + K + N - 2}",
f"| useful multiply-adds: {M * K * N}")

Output:

cycle 0: activations entering left = [1, 0, 0]
cycle 1: activations entering left = [4, 2, 0]
cycle 2: activations entering left = [0, 5, 3]
cycle 3: activations entering left = [0, 0, 6]
cycle 4: activations entering left = [0, 0, 0]
cycle 5: activations entering left = [0, 0, 0]
systolic result:
[[ 4  5  4]
[10 11 13]]
numpy X @ W:
[[ 4  5  4]
[10 11 13]]
match: True | cycles: 6 | useful multiply-adds: 18

Notice the overheads: the first ticks fill the array (the pipeline 'warms up') and the last ticks drain it. With a 256×256 array processing thousands of input rows, this fill-and-drain time is small compared with the steady state, where all 65,536 PEs do useful work every tick. With tiny inputs, most PEs sit idle — one reason TPUs like large batches.

The full journey of a TPU computation

Inside one layer, step by step

  1. Load a weight tile: A tile of W (for example 128×128) is moved into the matrix unit and held in the PEs.
  2. Stream activations: Thousands of activation rows are pumped through, skewed in time, while partial sums flow toward the output edge.
  3. Accumulate: Results from successive tiles along the inner dimension are added together in accumulators (TPU v1 kept these in dedicated on-chip accumulator memory).
  4. Apply the non-linearity: The vector unit applies activation functions and normalisation to the accumulated outputs.
  5. Next tile, next layer: Swap in the next weight tile (TPUs can preload it while the current one is in use) and repeat.

Why a TPU is so fast and power-efficient

  • Massive reuse. Each number fetched from memory is used across a whole row or column of PEs. Fewer memory accesses means less time waiting and much less energy.
  • Simple cells, packed densely. PEs have no instruction fetch, no branch prediction, no caches. Almost all of the chip's area and power goes to arithmetic and on-chip memory.
  • Low-precision maths. INT8 in v1, bfloat16 since v2 (and lower precisions in newer generations): smaller multipliers, more of them per chip, and half or less of the memory traffic of FP32.
  • Deterministic, compiled execution. XLA plans the whole computation ahead of time, so hardware does not need complex dynamic scheduling, and operations are fused to avoid round trips to memory.
  • Scale-out by design. Chips in a pod are linked directly, so huge models can be split across thousands of chips with predictable communication.

bfloat16 came from here bfloat16 ('brain floating point', named after Google Brain) keeps FP32's 8-bit exponent — the same range of magnitudes — with fewer mantissa bits. Training in it rarely needs special loss scaling. It is now supported by GPUs and other accelerators too.

Where TPUs are used

TPUs in practice
WhereWhat
Google productsSearch ranking, Translate, Photos, YouTube recommendations, speech recognition
Google's own modelsTraining and serving Gemini and other large models; earlier, AlphaGo and AlphaZero
Google Cloud customersCompanies and researchers rent TPU slices and pods for training and inference; Apple, for example, has reported training foundation models on TPUs
ResearchJAX-based research codebases; the TPU Research Cloud programme has given academics free access

Our shop's perspective If our appliance shop fine-tunes a model with JAX on Google Cloud, renting a TPU slice can be cost-effective for large, steady training jobs with fixed shapes. If we depend on CUDA-only libraries or need to run on our own servers, GPUs are the practical choice.

Limitations of a TPU

Where TPUs struggle Dynamic shapes (every new input shape can trigger a recompile), heavy data-dependent control flow, sparse or irregular operations that do not map to dense matrix tiles, small batches that leave the array mostly empty, and tensor sizes that do not align with the matrix unit (padding wastes compute). Custom low-level kernels are possible (for example with Pallas) but the ecosystem is smaller than CUDA's.

  • Availability: TPUs are offered through Google Cloud (and used internally), not sold as cards for our own servers.
  • Software lock-in: best support is for JAX and TensorFlow; PyTorch works through PyTorch/XLA, with some features and libraries lagging.
  • Debugging and profiling of compiled programs can feel less direct than eager-mode GPU code.
  • Not general-purpose: no graphics, little use outside dense tensor maths.

When not to choose a TPU: for small experiments that need many CUDA-specific libraries, for workloads with constantly changing shapes, or when data must stay on-premises.

Worked example, step by step

The simulation above produced the right answer. Now let us ask how well the tiny array was used. We stay with the same toy: a 3×3 weight-stationary array and X with M = 2 rows. Everything here is arithmetic on our own toy model; it shows the design idea and is not a measurement of any real TPU.

Scoring the 2 × 3 × 3 example

  1. Useful work: M × K × N = 2 × 3 × 3 = 18 multiply-accumulates, as the program printed.
  2. Time taken: M + K + N − 2 = 6 ticks: the staircase needs a few ticks to fill the array and a few to drain it.
  3. Capacity offered: 9 cells × 6 ticks = 54 cell-ticks. Only 18 of them did useful work: the array was 33% busy. The rest were cells waiting for data to arrive or leave.
  4. Memory reads with reuse: Each X value enters once (2 × 3 = 6 reads) and each weight is loaded once (9 reads): 15 reads in total.
  5. Memory reads without reuse: A design that fetched both operands for every multiply would need 2 × 18 = 36 reads. The array saved a factor of 2.4, even on this tiny job.
  6. Feed it more rows: With M input rows the busy share is M / (M + K + N − 2) = M / (M + 4). At M = 16 it is 80%; at M = 96 it is 96%. The fill and drain cost is paid once, so longer streams make it matter less.

This is the arithmetic behind two limits listed in this lesson. Small batches leave the array mostly empty, because fill and drain dominate. And sizes that do not line up with the array waste cells: in a toy with 128×128 tiles, a 130×128 weight matrix needs two tiles, and almost half of their cells would hold padding zeros. When a job runs slower than expected on this kind of hardware, batch size and layer sizes are the first two things to look at.

Practice: try it yourself

We will write a small calculator for our toy systolic array. It counts useful multiply-accumulates, ticks, how busy the cells are, and how many memory reads the reuse saves. Then it checks how well weight matrices of different sizes fit into square tiles.

practice_array_planner.py

import math
import numpy as np
def plan(M, K, N):
"""Cost of Y = X @ W (X is M x K, W is K x N) on a K x N toy array."""
macs = M * K * N                             # useful multiply-accumulates
naive_reads = 2 * macs                       # fetch both operands every time
reuse_reads = M * K + K * N                  # read each X and each W once
ticks = M + K + N - 2                        # fill + steady state + drain
busy = macs / (ticks * K * N)                # share of cell-ticks doing work
return macs, naive_reads, reuse_reads, ticks, busy
# Check the toy model against the lesson's 2 x 3 x 3 example.
X = np.array([[1, 2, 3], [4, 5, 6]])
W = np.array([[1, 0, 2], [0, 1, 1], [1, 1, 0]])
macs, naive, reuse, ticks, busy = plan(*X.shape, W.shape[1])
print(f"lesson example: {macs} MACs in {ticks} ticks, array busy {busy:.0%}")
print(f"  memory reads: {naive} without reuse, {reuse} with reuse")
print("more input rows through the same 3x3 array:")
for M in (1, 2, 16, 96):
macs, naive, reuse, ticks, busy = plan(M, 3, 3)
print(f"  M={M:3d}: ticks={ticks:3d}  busy={busy:4.0%}  "
f"reads saved={naive / reuse:.1f}x")
print("fitting a K x N weight matrix into square tiles of 128:")
for K, N in [(128, 128), (130, 128), (200, 200), (256, 256)]:
tiles = math.ceil(K / 128) * math.ceil(N / 128)
used = K * N / (tiles * 128 * 128)           # the rest is zero padding
print(f"  {K}x{N}: {tiles} tile(s), {used:.0%} of the cells hold real weights")

Output:

lesson example: 18 MACs in 6 ticks, array busy 33%
memory reads: 36 without reuse, 15 with reuse
more input rows through the same 3x3 array:
M=  1: ticks=  5  busy= 20%  reads saved=1.5x
M=  2: ticks=  6  busy= 33%  reads saved=2.4x
M= 16: ticks= 20  busy= 80%  reads saved=5.1x
M= 96: ticks=100  busy= 96%  reads saved=5.8x
fitting a K x N weight matrix into square tiles of 128:
128x128: 1 tile(s), 100% of the cells hold real weights
130x128: 2 tile(s), 51% of the cells hold real weights
200x200: 4 tile(s), 61% of the cells hold real weights
256x256: 4 tile(s), 100% of the cells hold real weights

Now change it:

  • In the second loop, call plan(M, 256, 256) with M in (1, 256, 4096). Predict the busy share for each. How many rows does a big array need before it is mostly full?
  • Add (129, 129) to the list of matrix sizes. Predict the number of tiles and the share of cells holding real weights.
  • Give X a third row, [7, 8, 9]. Predict the MACs, ticks and busy share printed on the 'lesson example' line.

Pause and think: In the run above, 'reads saved' rose from 1.5× to 5.8× as M grew, and it will never pass 6× for this 3×3 array. Why 6?

Without reuse we read 2·M·K·N = 18·M values. With reuse we read M·K + K·N = 3·M + 9. For large M the fixed 9 weight reads stop mattering, and the ratio tends to 18·M / 3·M = 6, which is 2 × N. Each X value is used N = 3 times after a single read, and each weight is reused for every row. A wider array (larger N) would save even more per value read.

Pause and think: In the toy tiling, a 130×128 matrix used 2 tiles with 51% real weights. A teammate suggests making the layer 128 wide instead of 130. What would we gain, and what must we check?

In the toy model the matrix would then fit one tile exactly: half the tiles and no padding, so no cells spend ticks multiplying zeros. What we must check is the model itself: a slightly narrower layer has fewer parameters, so we should confirm that accuracy does not suffer. The general habit is to pick sizes that line up with the hardware when the model does not care about the difference.

Key takeaways

  • A TPU is Google's ASIC for neural networks, built around matrix multiplication.
  • Its systolic array passes data between neighbouring cells so each value is fetched once and reused many times.
  • Reuse, simple cells, low precision (INT8, bfloat16) and compiled execution make it fast and power-efficient.
  • The XLA compiler turns JAX/TensorFlow/PyTorch code into static programs with fixed shapes.
  • TPUs excel at large, regular, dense workloads and struggle with dynamic shapes, irregular operations and small batches.

Key terms

  • TPU: Tensor Processing Unit: Google's custom chip for neural-network computation.
  • ASIC: Application-specific integrated circuit: a chip designed for one kind of task.
  • Systolic array: A grid of simple processing elements through which data flows rhythmically between neighbours.
  • Multiply-accumulate (MAC): The operation acc = acc + a × b, the building block of matrix multiplication.
  • Weight-stationary: A systolic dataflow where each cell holds a fixed weight while activations and partial sums move.
  • XLA: Accelerated Linear Algebra: the compiler that turns ML programs into optimised TPU (and GPU) code.
  • bfloat16: A 16-bit floating-point format with FP32's exponent range and fewer mantissa bits.

← 17.2 CUDA Kernels: Writing Parallel Code for NVIDIA GPUs · 17.4 Language Processing Units: A New Approach to LLM Inference →