LET'S BUILD GPT // FIELD MAP
← field map
PART 03 · SELF-ATTENTION42:13–62:00 · 20 min

The mathematical trick: from for-loops to softmax

Andrej Karpathy · Let's build GPT (2023) · part 03 of 8

Transcript: this part, with timestamps

TL;DR — Twenty minutes with no neural network in them at all. The question is mechanical: a token at position t needs to summarise everything at positions 0…t, for every t and every batch row, without a Python loop and without seeing the future. The answer is that a lower-triangular matrix of weights, multiplied against the sequence, is that summary — one matmul does all B×T summaries at once, and normalising the rows to sum to one turns the sums into averages. Karpathy then rewrites the same normalisation as softmax of a matrix that is zero below the diagonal and -inf above it, which looks like a pointless detour until you notice that those zeros are the only thing standing between this and self-attention. If you remember one thing: the weight matrix is data, not structure — every version in this part produces the same numbers, and the last one is written so the numbers can start being learned.

The lecture has just finished a bigram model that reaches a validation loss of about 2.5 and generates convincing gibberish, and its ceiling is obvious: each token is predicted from itself alone, so the model cannot use a single character of context. Everything after this point is about widening that window. But Karpathy does not open the Transformer paper here. He opens a scratch cell with a random 4×8×2 tensor in it and spends twenty minutes on an operation that would fit in a linear-algebra homework, because the operation is the load-bearing part — once you see that a masked, row-normalised matrix multiply is a batched weighted average over the past, self-attention in P04 is just the question of where the weights come from.

Outline, with timestamps

The problem, stated as tensors

The scratch example is deliberately smaller than the real model: B, T, C = 4, 8, 2 — four independent sequences, eight time steps each, two channels of information per step. In the actual script B is batch_size, T is block_size and C is the embedding width, but nothing in this part depends on those numbers, which is the point of shrinking them.

The eight vectors in a row are currently strangers. We want position 4 to be able to see positions 0 through 4 and nothing else — not 5, 6 or 7, because those are the answers we are training the model to guess. That asymmetry is the entire causal structure of a decoder-only language model, and here it shows up as a shape of a matrix rather than as a piece of theory.

The cheapest thing you could possibly do with that constraint is average. Position 4 becomes the mean of the five vectors at positions 0–4. Karpathy is blunt that this throws almost everything away:

just doing a sum or like an average is an extremely weak form of interaction — this communication is extremely lossy. We've lost a ton of information about the spatial arrangement of all those tokens. Karpathy, 44:15

He is right, and it does not matter, because the averaging is a placeholder for the aggregation. What we are actually building is the plumbing — the shape of the computation that moves information backwards along a sequence. The weights get interesting in P04; the lost spatial information comes back through the positional embedding at the end of this very part.

Version 1: the loop you would write first

# toy example, from the Colab notebook for the video
torch.manual_seed(1337)
B,T,C = 4,8,2 # batch, time, channels
x = torch.randn(B,T,C)

# version 1: for-loop it
xbow = torch.zeros((B,T,C))
for b in range(B):
    for t in range(T):
        xprev = x[b,:t+1] # (t,C)
        xbow[b,t] = torch.mean(xprev, 0)

xbow is short for bag of words — the term for a representation that keeps what appeared and discards the order. The slice x[b,:t+1] has shape (t+1, C), and torch.mean(xprev, 0) collapses the time axis, leaving a (C,) vector that gets written into xbow[b,t]. The +1 is the whole causal story: it is inclusive of the current token and exclusive of everything after it.

Print x[0] and xbow[0] side by side and the structure is visible without any linear algebra: row 0 of the two is identical (the average of one thing is that thing), row 1 of xbow is the midpoint of rows 0 and 1 of x, and row 7 is the mean of all eight. That hand-check at 46:49 is worth doing yourself, because it is the ground truth that the next two versions have to reproduce exactly.

It is also unusable. Two Python loops, B×T iterations, a fresh slice and a fresh reduction each time — for the real model that would be 64 × 256 = 16,384 sequential mean operations per forward pass, per layer, all of them tiny and all of them serialised through the interpreter.

Matrix multiply as weighted aggregation

The reframing takes three cells and no PyTorch beyond tril. Start with a 3×3 of ones times a 3×2 of small integers: every row of the output is the column-wise sum of the right operand, three times over, because every row of the left operand says "take all of you, equally".

torch.manual_seed(42)
a = torch.tril(torch.ones(3, 3))
a = a / torch.sum(a, 1, keepdim=True)
b = torch.randint(0,10,(3,2)).float()
c = a @ b

Two edits turn that into exactly what version 1 computes.

First, torch.tril keeps the lower triangle and zeroes everything above the diagonal, so a becomes [[1,0,0],[1,1,0],[1,1,1]]. A zero in a row of the left operand means that row of the right operand contributes nothing to that output row. Row 0 now sums a one-element prefix, row 1 a two-element prefix, row 2 all three. With b = [[2,7],[6,4],[6,5]] the product is [[2,7],[8,11],[14,16]] — running sums down the columns.

Second, row-normalise. a / torch.sum(a, 1, keepdim=True) divides each row by its own number of ones, giving [[1,0,0],[.5,.5,0],[⅓,⅓,⅓]], and the running sums become running averages: [[2,7],[4,5.5],[4.67,5.33]]. The keepdim=True is not cosmetic — it keeps the sum at shape (3,1) instead of (3,), which is what lets it broadcast across columns rather than across rows. Drop it and you silently normalise the wrong axis.

This is the sentence to carry out of the part: a matrix multiply is a weighted aggregation, and the left operand is the table of weights. Whatever you can express as "each output is some weighted combination of the inputs" is a matmul waiting to happen; the lower-triangular shape is just the particular weighting that says only look backwards. Every later version — one head, six heads, six layers — changes what fills that table and never changes the operation.

Version 2: one batched matmul

# version 2: using matrix multiply for a weighted aggregation
wei = torch.tril(torch.ones(T, T))
wei = wei / wei.sum(1, keepdim=True)
xbow2 = wei @ x # (T, T) @ (B, T, C) ----> (B, T, C)
torch.allclose(xbow, xbow2)

Here wei — weights, spelled that way throughout the repo and heard as "way" in the captions — is (T,T) = (8,8), and x is (B,T,C) = (4,8,2). The shapes do not match, and PyTorch does not complain: for @, any dimensions before the last two are batch dimensions, so a missing one is broadcast into place. wei is treated as (1,8,8), stretched to (4,8,8), and the same 8×8 matrix is applied independently to each of the four sequences. One kernel launch replaces 32 Python iterations, and at real scale it replaces 16,384 of them.

tensorshape in the toyshape in gpt.pywhat it holds
x(4, 8, 2)(64, 256, 384)the sequence, C numbers per token
wei(8, 8)(64, 256, 256)how much of row j flows into row i
tril(8, 8)(256, 256)the causal mask, a constant buffer
wei @ x(4, 8, 2)(64, 256, 384)same shape in, same shape out — always

Note the last row. The aggregation never changes the shape of the sequence; it only changes what is written at each position. That invariance is why you can stack six of these blocks in P06 without any bookkeeping.

Version 3: softmax over a masked matrix

# version 3: use Softmax
tril = torch.tril(torch.ones(T, T))
wei = torch.zeros((T,T))
wei = wei.masked_fill(tril == 0, float('-inf'))
wei = F.softmax(wei, dim=-1)
xbow3 = wei @ x
torch.allclose(xbow, xbow3)

Read it as three moves. wei starts as all zeros — an affinity of zero between every pair of positions. masked_fill writes -inf wherever tril is zero, i.e. everywhere above the diagonal, which is every (query position, future position) pair. Then softmax along dim=-1 exponentiates and normalises within each row: exp(0) = 1 in the allowed cells, exp(-inf) = 0 in the masked ones, and dividing by the row sum gives row t a uniform 1/(t+1) over its prefix and hard zero elsewhere. Bit-for-bit the same matrix as version 2's hand-normalisation, arrived at by a route that happens to accept arbitrary numbers where the zeros are.

That is the only reason for the detour, and Karpathy previews it at 56:44: those zeros are placeholders for how interesting each past token is to the current one. Once they are computed from the data instead of hard-coded, softmax still produces a valid set of mixing weights — non-negative, summing to one per row — and -inf still guarantees that the future contributes exactly nothing, because it is the one value that survives exponentiation as a true zero rather than a small number. Writing -1e9 instead would leak, slightly, and forever.

the elements here in the lower triangular part are telling you how much of each element fuses into this position. Karpathy, 58:17

Cleanup: a bottleneck and a head

With the trick in hand, the script needs two structural changes before attention can be dropped into it. Both are small and both are load-bearing.

In the bigram model the embedding table was (vocab_size, vocab_size) — a token index went straight in and logits came straight out. Karpathy inserts a level of indirection at 58:26: the table becomes (vocab_size, n_embd) with n_embd = 32, and a new nn.Linear(n_embd, vocab_size) called lm_head maps embeddings to logits. Nothing is gained yet — it is one extra linear layer on a model with no context, and he says so — but there is now a 32-wide vector per token that is not the output distribution, which is the space every later mechanism operates in. (The 32 was suggested by Copilot on screen and kept; it survives as n_embd = 384 in the final script.)

Then the piece that pays back the "extremely lossy" complaint from twenty minutes earlier. A second table, position_embedding_table, of shape (block_size, n_embd), is indexed by torch.arange(T) — the integers 0…T-1 — producing a (T, C) tensor that is added to the (B, T, C) token embeddings. The addition broadcasts: the position tensor gets a leading dimension of 1 and is stretched across the batch, so every sequence in the batch is offset by the same position vectors. Positions are learned here, one free vector per slot, not the sinusoids of the original paper.

At this exact moment the positional information does almost nothing. The logits at step t still depend only on the token at t and the position t, so all the model can learn from it is a position-dependent bias on the output distribution — mildly useful for "line 1 tends to start with a capital", useless for language. It is here because attention has no notion of order (P05, note 2), and by the time it can read x, the order had better already be encoded in it.

The code at the end of this part

The permanent residue in the final script is the two embedding tables and the head, plus the first four lines of the forward pass:

        # each token directly reads off the logits for the next token from a lookup table
        self.token_embedding_table = nn.Embedding(vocab_size, n_embd)
        self.position_embedding_table = nn.Embedding(block_size, n_embd)
        self.blocks = nn.Sequential(*[Block(n_embd, n_head=n_head) for _ in range(n_layer)])
        self.ln_f = nn.LayerNorm(n_embd) # final layer norm
        self.lm_head = nn.Linear(n_embd, vocab_size)
    def forward(self, idx, targets=None):
        B, T = idx.shape

        # idx and targets are both (B,T) tensor of integers
        tok_emb = self.token_embedding_table(idx) # (B,T,C)
        pos_emb = self.position_embedding_table(torch.arange(T, device=device)) # (T,C)
        x = tok_emb + pos_emb # (B,T,C)
        x = self.blocks(x) # (B,T,C)
        x = self.ln_f(x) # (B,T,C)
        logits = self.lm_head(x) # (B,T,vocab_size)

Read those against the video with one substitution: at 62:00 the self.blocks and self.ln_f lines do not exist yet — x goes straight from tok_emb + pos_emb into lm_head. Everything else is already in its final form. gpt.py#L143–L147 is the constructor; gpt.py#L160–L169 is the forward. B, T = idx.shape is unpacked explicitly because T is needed for arange and is not always block_size — during generation the context starts at length 1 and grows, which is exactly why generate later has to crop with idx[:, -block_size:] (L185): index the position table past block_size and it throws.

And the trick itself, three lines of it, survives verbatim inside the attention head:

        self.register_buffer('tril', torch.tril(torch.ones(block_size, block_size)))
        ...
        wei = wei.masked_fill(self.tril[:T, :T] == 0, float('-inf')) # (B, T, T)
        wei = F.softmax(wei, dim=-1) # (B, T, T)

L72, L84–L85. Three details that only make sense having watched this part: register_buffer rather than a plain attribute or a Parameter, because tril is a constant that must follow the model to the GPU but must never receive a gradient; self.tril[:T, :T] rather than self.tril, because the buffer is allocated at full block_size and cropped to whatever T the current input has; and wei arriving as (B,T,T) instead of the toy's (T,T), because in P04 each batch row computes its own affinities. The mask broadcasts across the batch; the affinities do not.

Where people get stuck

Go deeper, verified

Exercises

  1. Break the mask on purposecode — in the Colab, replace tril with a banded mask that lets each token see only itself and the previous three (hint: torch.tril(ones) - torch.tril(ones, -4)). Re-run version 3, then write the matching double for-loop and check the two against each other with allclose. A good answer notes which rows of the new wei no longer sum to one before the softmax and why the softmax makes that irrelevant, and states what the mask means as a model: a fixed four-token context window, i.e. sliding-window attention.
  2. Make the affinities non-uniform by handcode — keep version 3 but initialise wei = torch.randn(T,T) instead of zeros before the masked_fill. Print wei after the softmax and confirm each row is non-negative, sums to one, and is zero above the diagonal. Then multiply the pre-softmax matrix by 10 and by 0.1 and print again. A good answer describes the two limits — one-hot "hard attention" as the scale grows, uniform averaging as it shrinks — and connects that to why P05 spends a whole note on dividing by sqrt(head_size).
  3. Beat the matmul, then explain why you shouldn'tcode — implement the same prefix average a third way with torch.cumsum(x, dim=1) divided by torch.arange(1, T+1) (get the broadcast right: it needs shape (1,T,1)). Verify it matches xbow, and time all three at T = 1024. A good answer observes that cumsum is O(T) against the matmul's O(T²) and wins outright — and then explains why the lecture takes the slower road anyway: cumsum can only ever compute a uniform average, whereas the matmul accepts any weight matrix, and the weights are the entire point.
Next: P04 The crux: self-attention · Back to the map.