The mathematical trick: from for-loops to softmax
Transcript: this part, with timestamps
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
- 42:13 — The toy tensor: B, T, C = 4, 8, 2, eight tokens per row that currently do not talk to each other.
- 43:44 — The constraint: information flows from the past to the present only, because the future is what we're predicting.
- 44:46 — Version 1: the honest double for-loop that builds xbow, a "bag of words" average, then read by hand at 46:49.
- 47:11 — The trick, in miniature: a 3×3 of ones times a 3×2 gives column sums, repeated.
- 48:25 — torch.tril zeroes the upper triangle, so each output row sums a different-length prefix.
- 50:29 — Dividing each row by its own sum turns the sums into running averages.
- 51:54 — Version 2: one (T,T) @ (B,T,C) batched matmul replaces both loops; allclose confirms it.
- 54:42 — Version 3: start from zeros, masked_fill the future with -inf, take a softmax.
- 56:44 — Why bother: those zeros are affinities, and affinities are about to become data-dependent.
- 58:26 — Cleanup: an n_embd = 32 bottleneck plus an lm_head linear layer, so the embedding is no longer the logits.
- 60:18 — Positional encoding: a second embedding table indexed by position, added to the token embedding.
- 61:20 — The broadcast that makes (T,C) add cleanly to (B,T,C), and why it does nothing useful yet.
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.
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.
| tensor | shape in the toy | shape in gpt.py | what 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
- Which direction the mask points. Karpathy says it backwards on tape at 57:00 — "tokens from the past cannot communicate" — and corrects it in the video description: it is the future that cannot communicate. Fix the direction by reading the indices, not the words: wei[i,j] is how much position j contributes to position i, tril keeps j <= i, and each row of the matrix is one query position looking back down its own prefix. Row-major, past-only, current token included.
- Why (T,T) @ (B,T,C) is legal. It looks like a shape error and is not. torch.matmul treats everything before the final two dimensions as batch and broadcasts it, so the 8×8 is promoted to (1,8,8) and applied to each of the four sequences. The same rule quietly fires again one section later on (B,T,C) + (T,C). If a shape puzzle in this lecture confuses you, it is almost always broadcasting — and the debugging move is to print the shapes rather than reason about them.
- Nothing here is learned. wei in version 3 is torch.zeros, a literal, not a parameter — running these cells trains nothing and changes no loss. The zeros exist to be replaced. It is a fair criticism that the softmax rewrite is unmotivated until P04, and that is exactly the shape of the part: version 3 is a refactor written in advance of the feature it enables.
- torch.allclose comes back False. On some PyTorch builds version 2 and version 1 disagree in the last bits, because a matmul accumulates in a different order than torch.mean and float addition is not associative. Check with (xbow - xbow2).abs().max(); if it is around 1e-8 or smaller, the code is right and the tolerance is wrong. Pass atol=1e-6 and move on.
Go deeper, verified
- Colab notebook for the lecture — Andrej Karpathy (2023) · the four versions live in scratch cells here and nowhere in the repo; run them and print the intermediate matrices, which is the fastest way to make this part stick.
- ng-video-lecture · Head.forward — Andrej Karpathy (2023) · fifteen lines where the trick ends up, with the shape of every intermediate in a trailing comment. Worth reading now even though half of it belongs to P04.
- Attention Is All You Need — Vaswani et al. (2017) · §3.2.3 is the paragraph this part implements: illegal connections in the decoder are masked by setting them to −∞ before the softmax. One sentence in the paper, twenty minutes on video.
- PyTorch broadcasting semantics — PyTorch docs · settles both of this part's shape surprises, and most of the ones in P04 and P06 too.
- nanoGPT · model.py — Andrej Karpathy (2023) · how the same mask is done in production: a bias buffer of the same triangular ones, with a fast path that hands the causality to PyTorch's fused kernel instead.
- F.scaled_dot_product_attention — PyTorch docs · field map extra · the modern one-liner. Its is_causal=True flag is precisely the tril/masked_fill pair you just wrote by hand.
- FlashAttention — Dao et al. (2022) · field map extra · the sequel to this part's central move. Materialising a T×T matrix is O(T²) in memory, which is fine at T = 8 and fatal at T = 128k; FlashAttention computes the same result while never writing that matrix down.
- RoFormer (RoPE) — Su et al. (2021) · field map extra · what replaced the learned position table added at 60:18 in essentially every model since 2022, by rotating queries and keys instead of adding a vector to the input.
Exercises
- 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.
- 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).
- 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.