1Day 1

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.

The problem

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.

Intuition

Queries, keys and values

Each position turns its vector into three new vectors with three learned matrices:

Position ii 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 ii 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.

Playground
x: [1, 8, 16]

Eight positions, each a vector of 16 numbers. In a trained model these come from the embedding table. Here they are random draws from torch.randn with seed 0.

x8 × 16
0123456789101112131415t0t1t2t3t4t5t6t7
Step 1 of 8
Hover a row in the scores, mask or softmax step to highlight what one position reads. Positions t0 to t7 are sequence positions, not words. The inputs are random, so the patterns carry no meaning yet.

What happened?

The math

One formula, read left to right

The whole layer for one head fits in one line:

Attention⁡(Q,K,V)=softmax⁡ ⁣(QK⊤d+M)V\operatorname{Attention}(Q, K, V) = \operatorname{softmax}\!\left(\frac{QK^\top}{\sqrt{d}} + M\right) V

QQ, KK and VV hold one row per position. dd is the head width, 8 here. MM is the mask:

Mij={0if j≤i−∞if j>iM_{ij} = \begin{cases} 0 & \text{if } j \le i \\ -\infty & \text{if } j > i \end{cases}

The division by d\sqrt{d} 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 dd such products and has variance dd. At d=64d = 64 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 d\sqrt{d} brings the variance back to 1.

The shapes for the playground's input, with batch B=1B = 1, length T=8T = 8, width C=16C = 16 and 2 heads:

StepShape
input xx[1, 8, 16]
q,k,vq, k, v after projection[1, 8, 16] each
split into heads[1, 2, 8, 8] each
scores qk⊤/8qk^\top / \sqrt{8}[1, 2, 8, 8]
weights after mask and softmax[1, 2, 8, 8]
weights times vv[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

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 ii is the token at position i+1i + 1, and that token is sitting right there in the input.

Without the mask, position ii could attend to position i+1i + 1 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 pp must leave every output row before pp 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:

Check
Mask
t00
t10
t20
t30
t40
t50
t60
t71.7e-1
Rows t0 to t6 changed by exactly 0. No position read the future.

Each bar is the largest change in that row of the attention output after adding a random vector to x[t7].

Each bar is one output row. Green values are exactly zero. Switch the mask to None to see what the model would learn to exploit.

What happened?

In octlm

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:

octlm/model.pyline 293
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)) @ value

masked_fill(~mask, float("-inf")) is the +M+M 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:

octlm/model.pyline 178
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)
Multi-head

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 cost

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.

Playground
Q64×32×Kᵀ32×64→scores, masked64×64×V64×32→out64×32
Score entries4,096T², the orange square
Q, K, V entries6,1443 × T × d
Multiply-adds262,144QKᵀ plus weights × V

Q, K, V and the output grow in a straight line with T. The score square grows with T². Double the length and the square quadruples, whatever the head width. The gray triangle above the diagonal is the causal mask: positions it covers are computed and then thrown away.

The matrices of one head, drawn to scale. Drag the sequence length and watch the score square outgrow everything else.
Skipped

What we did not build, and why