Parallelizing the Transformer, Masking the Future, and the Language Modeling Head
Last post built self-attention and the transformer block, but left two promises unkept: I showed the computation one token at a time (so where's the famous parallelism?), and I never showed how any of it actually predicts a word. This post closes both gaps. By the end, you'll know how the whole attention computation collapses into a few big matrix multiplies that run in parallel, why a causal language model has to mask out the future and how a triangle of negative infinity does it, where the input vectors actually come from (token embeddings plus a second thing we've been ignoring: position), and how the language modeling head turns the top-of-stack vector into a probability distribution over the next word. That last piece is the literal proof of the claim that an LLM "just predicts the next token." We'll finish with how these models scale and how you fine-tune one without melting your GPU.
Recap: Where We Left Off
Quick reset. A transformer block takes a token's d-dimensional vector, runs it through layer norm, multi-head attention, another layer norm, and a feedforward layer (with residual connections threading the original vector through), and emits a d-dimensional vector. Same shape in, same shape out, so you stack these blocks, 12 to 96 or more. Only the attention step looks at other tokens; everything else refines a token in place.
We did all of that one token, one residual stream at a time. But here's what we need to emphasize: the computation for each token is independent of the others (except inside attention). That independence is what lets us stop looping and start multiplying matrices.
Doing It All at Once: The Matrix X
Instead of processing token i, then token i+1, then i+2, we pack the entire input sequence into one matrix. Call it , of shape — N rows (one per token), each row the d-dimensional embedding of that token. For real models N (the context window) runs from 1K to 32K tokens and beyond.
Now the per-token query/key/value projections become single matrix multiplies over the whole sequence at once:
With of shape and of shape , this gives of shape and of shape — every token's query, key, and value computed in one shot.
Now the satisfying part. To get all the query-key comparisons (every token scored against every other token), you just multiply by :
Cell of this matrix is , the relevance of token j to token i. The entire attention score table, in a single matrix multiplication. Scale it, softmax it, multiply by , and you have the attention output for every token at once:
That mask in there is new. It's the price of doing everything at once, and it's worth understanding properly.
Masking the Future
The matrix computes a score for every query against every key, including keys that come after the query in the sequence. For a causal language model, that's cheating.
Think about what the model is trained to do: predict the next word. If, while computing the representation of token 3, you let it attend to tokens 4, 5, 6… you've handed it the answer. Guessing the next word is pretty easy if you already know it.
So we mask out the future. In the score matrix, every cell where the key is ahead of the query (the upper triangle) gets set to :
Why negative infinity? Because the very next step is a softmax, and . Those future positions get exactly zero attention weight. The model is mathematically blindfolded to anything it hasn't reached yet.
One Catch: Attention Is Quadratic
That matrix is . Which means the cost of attention grows with the square of the sequence length; you're computing a dot product between every pair of tokens.
Double the context window, and you quadruple the attention computation. This is why long context is expensive, why running attention over an entire novel is hard, and why "rooms full of GPUs" is not a figure of speech. Modern models still push contexts into the tens of thousands of tokens, but the quadratic cost is the wall everyone is pushing against.
The Parallel Block, and the Two Big Differences
With attention parallelized, the whole transformer block has a clean matrix form (prenorm — layer norm before each sublayer):
Both and are : the full sequence in, the full sequence out, so blocks stack exactly as before. For the first block, is the input embeddings; for block k, it's the output of block k−1.
Look at how far we've come from RNNs. The transformer introduces exactly two new ideas relative to a recurrent net:
- Attention replaces recurrence. No hidden state is passed step to step; instead, every token directly attends to every (prior) token.
- Parallel computation replaces iteration. No waiting for token i−1 to finish before starting token i; the whole sequence goes through the matrix multiplications together.
Everything else (embeddings, softmax, predicting the next word) we've seen since the n-gram days. Those two ideas are all that's actually new.
X (input sequence) → [N × d], one row per token. W^Q, W^K → [d × dₖ]; W^V — [d × dᵥ]. Q, K → [N × dₖ]; V → [N × dᵥ]. QKᵀ (all query-key scores) → [N × N]. Attention output A → [N × d]; same shape as X, so blocks stack. Logits u (one score per vocab word) → [1 × |V|].Quick reference: the shapes
What Is X, Really? Token and Position Embeddings
I've been waving my hands about "the input embeddings." Time to be honest about what actually contains, because it's not just token embeddings. Each row is the sum of two embeddings.
Token embeddings
The familiar part. There's an embedding matrix of shape — one row per vocabulary token. To get the input, you tokenize the string (with byte-pair encoding, which can split into subword pieces), convert each token to its vocabulary index, and select the matching rows of . So "Thanks for all the" might become indices , and you pull rows 5, 4000, 10532, 2224.
Position embeddings
One thing here is easy to miss: Because we process all tokens in parallel, the model has no inherent sense of order. Nothing in the matrix multiplication says token 3 came before token 4. We have to inject that information explicitly.
So alongside the token embedding, we add a position embedding: a vector that says "this is position 3." The simplest scheme, absolute position, keeps a learned embedding matrix with one randomly initialized vector per position, learned during training just like word embeddings. Just as the embedding for fish captures what fish means, the embedding for position 3 captures something about what tends to appear in position 3.
The final input is their sum:
Both are , so the sum is too. That composite vector (word meaning + position) is what actually enters the first transformer block.
The Language Modeling Head
We have a stack of transformer blocks turning input embeddings into rich contextual representations. But a pile of d-dimensional vectors isn't a prediction. The language modeling head is the final circuit that turns the top of the stack into an actual next-word distribution.
The recipe is short. Take the output of the last token from the last layer, call it , a single vector. Don't underestimate this vector: it has been through every layer of attention and feedforward, so it's a deep contextual summary of everything seen so far. If the model has done its job, it encodes exactly what's needed to guess what comes next.
Now project it to a score for every word in the vocabulary. We multiply by an unembedding matrix, and there's a nice trick here called weight tying: instead of learning a fresh matrix, we reuse the transpose of the input embedding matrix, , of shape . At the input, maps a word to an embedding; at the output, maps an embedding back to scores over words. Same weights, run in reverse, hence "unembedding."
The vector holds the logits — raw, unnormalized scores, one per vocabulary word ( ). The softmax turns them into a probability distribution over the whole vocabulary. Then you sample a word from — greedily (take the most probable), or with top-k / top-p sampling for more diversity (which we covered back in the LLM post).
One bit of terminology: a transformer used this way (left-to-right, masked, predicting the next token) is called a decoder-only model. It's the decoder half of the original encoder-decoder transformer, repurposed.
The Full Picture
Stand back and trace one token through the finished machine:
- Input encoding — token embedding + position embedding → .
- L transformer blocks — each applying masked multi-head attention and a feedforward layer over the residual stream, all N tokens in parallel, the representation getting richer at every layer.
- Language modeling head — take the final layer's last-token vector, unembed with , softmax into a distribution, sample the next word.
That's a transformer language model, end to end. We started this course counting n-grams; we're now generating text with a 96-layer attention network. Same task the whole way (predict the next word), just a vastly better predictor.
Dealing with Scale
A quick tour of what changes when these models get big, because "big" is doing a lot of work in "large language model."
Scaling laws. Performance is governed mostly by three things: model size (parameters, not counting embeddings), dataset size, and compute. The empirical finding is that the loss falls as a power law in each one — smooth, predictable curves. That's useful in a concrete way: you can look at the early part of a training run and predict what the loss would be with more data or a bigger model, before spending the money.
The non-embedding parameter count has a tidy approximation:
Plug in GPT-3's 96 layers and d = 12,288 and you get roughly 175 billion parameters. That formula is why "add layers, widen the model" translates so directly into parameter counts.
KV cache. Training parallelizes beautifully, but generation is inherently one-token-at-a-time. When you generate token i, you'd otherwise recompute the key and value vectors for all the earlier tokens — wasteful, since you already computed them. The fix is to cache those key/value vectors in memory and reuse them. A small idea that makes inference much cheaper.
Fine-Tuning Without Melting Your GPU: LoRA
One last practical piece. Suppose you have a pretrained model and want to specialize it for your domain. Full fine-tuning means updating all the parameters, backpropagating through every one of those billions of weights, every step. For a large model that's expensive in compute and memory, and slow.
Parameter-efficient fine-tuning sidesteps this by updating only a small subset. The most popular method is LoRA (Low-Rank Adaptation), and the idea is elegant. A weight matrix (say ) would normally get an update of the same large size. Instead, LoRA freezes and writes the update as a product of two skinny matrices:
where is , is , and the rank is tiny — often 1 or 2. You only train and , a few thousand parameters instead of millions. Since is the same shape as , you can just add it in — no extra cost at inference, and you can swap different LoRA modules in and out for different domains. In its original form it was applied to the attention matrices ( ).
Causal mask: the −∞ upper triangle that stops a token from attending to future tokens. Position embedding: a learned vector added to each token embedding to encode its place in the sequence. Logits: the raw, unnormalized scores over the vocabulary, before softmax. Unembedding (Eᵀ): the matrix that maps a final-layer vector back to vocabulary scores; weight-tied to the input embedding E. Weight tying: reusing the same matrix (here E and its transpose) in two places. Decoder-only: a causal, masked, left-to-right transformer language model (the decoder half of the original encoder-decoder). KV cache: stored key/value vectors reused during generation so they aren't recomputed each step. Scaling laws: the power-law relationship between loss and model size / data / compute. LoRA: low-rank fine-tuning: freeze W, train a small A·B instead.New terms in this chapter, at a glance
What You Now Have
Seven things from this lecture:
Parallel attention via the matrix X. Pack the sequence into of shape , compute (and K, V) in one shot, and get every query-key score at once as the matrix .
Causal masking. Set the upper triangle of to so the softmax zeroes out the future. This is how a model can be both parallel and honest about not seeing ahead.
Attention is quadratic in sequence length, a dot product per pair of tokens, which is why long context is expensive.
The two transformer ideas vs. RNNs: attention (replaces recurrence) and parallel computation (replaces iteration). Everything else is older machinery.
The input is token + position embeddings. Parallel processing erases order, so we add a learned position embedding to each token embedding. Absolute positions undertrain at large indices, which motivates sinusoidal and relative schemes.
The language modeling head. Take the last token's final-layer vector , multiply by the unembedding matrix (weight-tied to ) to get logits, softmax into a distribution over the vocabulary, sample. The exact mechanism behind "predicting the next token."
Scale and fine-tuning. Loss follows power-law scaling laws in parameters/data/compute ( , ~175B for GPT-3); the KV cache speeds inference; and LoRA fine-tunes cheaply by freezing and training a tiny low-rank instead.






Top comments (0)
Some comments may only be visible to logged-in visitors. Sign in to view all comments.