Causal attention, by hand
How each position in a sequence decides which earlier positions to read, traced step by step through octlm's attention layer with the weights PyTorch actually produced.
A language model reads a sequence of tokens and predicts the next one at every position. Most of the layers in a decoder work on one position at a time. The feed-forward layer, the norms and the output head each see a single vector and nothing else.
Attention is the only layer where positions exchange information. If the model is to predict what follows "The cat sat because it", the vector at "it" has to learn something about "cat". Attention is how it gets there.
This post walks through octlm's attention layer from input to output, on a real module with 16-dimensional vectors and two heads. Then it checks the one property a decoder cannot live without.
Every position needs to read the positions before it
Older models had two ways to carry context forward. An n-gram model looks at a fixed window, so a word four tokens back is invisible to a 3-gram. A recurrent network squeezes the whole history into one vector that it updates token by token, so distant tokens fade.
Attention does neither. Each position looks at every earlier position directly and decides how much of each to take. The distance between "it" and "cat" does not matter. What matters is whether the model learned that they are related.
Queries, keys and values
Each position turns its vector into three new vectors with three learned matrices:
- A query describes what this position is looking for.
- A key describes what this position can offer to others.
- A value is the information this position passes on if someone reads it.
Position compares its query with the key of every position it may see. A large dot product means a good match. The matches become weights, and the output at is the weighted average of the values.
Step through it below. The numbers are real. scripts/export.py builds octlm's CausalSelfAttention with torch.manual_seed(0), feeds it random inputs, and saves every weight. The page recomputes each step in TypeScript, and a test checks the final output against PyTorch's.
What happened?
- In the Softmax step, row t0 has a weight of 1.00 on t0 and nothing else. The first position can only read itself.
- The top-right triangle is empty from the Mask step onward. No position has any weight on a position after it.
- Every row of the weights sums to 1, so each output row is an average of value rows, never a sum that grows with length.
- Head 1 and head 2 produce different weights from the same input, because each head has its own slice of the query and key matrices.
One formula, read left to right
The whole layer for one head fits in one line:
, and hold one row per position. is the head width, 8 here. is the mask:
The division by has a concrete reason. If the entries of a query and a key are independent with variance 1, their dot product is a sum of such products and has variance . At the scores would have a standard deviation of 8. Softmax over numbers that far apart puts almost all the weight on one position, and its gradient for the others nearly vanishes. Dividing by brings the variance back to 1.
The shapes for the playground's input, with batch , length , width and 2 heads:
| Step | Shape |
|---|---|
| input | [1, 8, 16] |
| after projection | [1, 8, 16] each |
| split into heads | [1, 2, 8, 8] each |
| scores | [1, 2, 8, 8] |
| weights after mask and softmax | [1, 2, 8, 8] |
| weights times | [1, 2, 8, 8] |
| heads joined, output projection | [1, 8, 16] |
The scores are [T, T] per head. That square is the cost of attention, and it comes back at the end of this post.
The mask is what separates reading from cheating
Training does not generate one token at a time. It feeds the whole sequence in at once and computes a prediction at every position in parallel. The target at position is the token at position , and that token is sitting right there in the input.
Without the mask, position could attend to position and copy the answer. The training loss would fall fast and the model would learn nothing it can use at generation time, when the next token does not exist yet.
So the property to check is exact. Changing the token at position must leave every output row before unchanged. The rows must be the same numbers, not merely close. octlm's Day 1 test does this. It changes the last input token and confirms that every earlier logit is byte-identical.
Try it on the attention layer alone:
What happened?
- With the causal mask, every row before the edited position changes by exactly 0. Those rows never read the edited value, so the arithmetic that produces them is the same arithmetic as before.
- Rows after the edited position do change. They are allowed to read it.
- With no mask, every row changes. Position t0 now depends on t7, which in training would be information from the future.
The code
The Day 1 attention path is written out by hand so every step is visible. This is the method that computes the weights:
def _naive(self, query: Tensor, key: Tensor, value: Tensor, mask: Tensor) -> Tensor:
groups = self.n_heads // key.shape[1]
if groups > 1:
key = key.repeat_interleave(groups, dim=1)
value = value.repeat_interleave(groups, dim=1)
scores = query @ key.transpose(-2, -1) / math.sqrt(query.shape[-1])
scores = scores.masked_fill(~mask, float("-inf"))
return self.dropout(scores.softmax(dim=-1)) @ valuemasked_fill(~mask, float("-inf")) is the in the formula. softmax(dim=-1) normalizes each row. The repeat_interleave branch does nothing on Day 1, where every query head has its own key and value head. It exists for grouped-query attention, which arrived on Day 2.
The mask comes from one function. Later experiments reuse it to build sliding-window and strided masks:
def attention_mask(
length: int, window: int, stride: int, device: torch.device | None = None
) -> Tensor:
"""Causal mask, optionally narrowed to a sliding window with strided global columns."""
rows = torch.arange(length, device=device).unsqueeze(1)
columns = torch.arange(length, device=device).unsqueeze(0)
distance = rows - columns
allowed = distance >= 0
if window:
local = distance < window
if stride:
local = local | (columns % stride == 0)
allowed = allowed & local
return allowed.view(1, 1, length, length)Why two heads instead of one wide one
Two heads of width 8 use the same parameters as one head of width 16. The difference is in the weights. One head produces one weight row per position, so each position can take one mixture of the others. Two heads produce two independent mixtures. In a trained model, one head might track the previous token while another tracks the subject of the sentence.
The heads cost nothing extra because the split is a reshape. [1, 8, 16] becomes [1, 2, 8, 8] without copying any numbers. The output projection mixes the heads back together.
The square that grows
The scores matrix is [T, T] for every head. Double the length and it has four times the entries. On Day 2 we measured this on the laptop CPU with PyTorch's plain math backend, which builds that matrix in memory. At 1,024 tokens a forward pass grew resident memory by 107 MiB. At 8,192 tokens it grew by 4,928 MiB. Those numbers come from runs/day2-exp014-sdpa.jsonl.
Day 2 replaced this path with PyTorch's fused attention kernel, which never stores the full square. The Day 1 path stays in the code as the reference that the fused one is checked against.
What we did not build, and why
- Pictures of a trained model's attention. The Day 1 model had 541,952 parameters and trained for 200 steps on one block. Its attention patterns would show noise. The Day 4 model is large enough and trained long enough to be worth looking at.
- Dropout on the attention weights. The layer supports it, but every config so far sets
dropout = 0.0. At these data sizes the models are undertrained, not overfit. - A KV cache. Generation recomputes the whole sequence for every new token. Caching keys and values is EXP-071 on Day 5.
- Cross-attention. An encoder-decoder model attends from one sequence to another. octlm is decoder-only, so every attention layer is self-attention.