USAAIO Lesson 85, from Phase 3. It covers the three sublayers of the transformer decoder block - causal masked self-attention, using an upper-triangular -inf mask; cross-attention, where Q comes from the decoder while K and V come from the encoder; and the feed-forward network - along with autoregressive generation. Every attention weight, shape, and row sum was verified with torch 2.7.1+cpu at seed 42. The lesson runs to 27 slides.
Subject: Machine Learning · 56 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 85 · Phase 3
Three sublayers, one autoregressive loop: causal masked self-attention, cross-attention, and FFN. How the decoder reads the encoder and generates tokens one at a time — verified in PyTorch from scratch.
Objectives
Warm-up
Discussion prompt
Before we open Lesson 85: Transformer Decoder Block: without looking back, what was the main idea of Transformer Encoder for Classification, and what could you do by the end of it that you could not do before?
Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.
Answer:
stacking N encoder blocks, sinusoidal positional encoding, [CLS] token pooling, padding-mask construction, and the full encoder classifier (AdamW + linear warmup + cosine decay + label smoothing). Build TransformerEncoderClassifier from scratch, verify permutation invariance without positional encoding, and confirm all encoder-block parameter counts.
Section
Part 1 of 4
Concept
An encoder block has two sublayers: self-attention and FFN. A decoder block adds a third between them: cross-attention. All three use residual connections and LayerNorm (post-norm in the original paper).
\[ \text{DecoderBlock}(x, z) = \text{LN}\bigl(\text{FFN}(\cdot) + \cdot\bigr) \circ \text{LN}\bigl(\text{CrossAttn}(\cdot, z) + \cdot\bigr) \circ \text{LN}\bigl(\text{MaskedSelfAttn}(x) + x\bigr) \]
Matching
Match the pairs
From Decoder block anatomy — match each one to what it actually does. The descriptions have been shuffled.
Why: ① Masked self-attention, ② Cross-attention, ③ FFN are easy to tell apart while they are sitting next to their descriptions and much harder afterwards, which is what this checks.
Section
Part 2 of 4
Concept
During training the full target sentence is fed in parallel. Without masking, position 2 would attend to position 5 — information that should not exist yet. Masking enforces the autoregressive property: each output depends only on previous outputs.
\[ \text{Causal mask}_{ij} = \begin{cases} 0 & j \le i \\ -\infty & j > i \end{cases} \]
For T=5, torch.triu(ones(5,5), diagonal=1) gives the upper-triangular pattern. Those positions get -inf added to the raw attention logits — before softmax.
Counterexample
Discussion prompt
For T=5, torch.triu(ones(5,5), diagonal=1) gives the upper-triangular pattern. Those positions get -inf added to the raw attention logits — before softmax.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Concept
| token | can see pos 0 | can see pos 1 | can see pos 2 | can see pos 3 | can see pos 4 |
|---|---|---|---|---|---|
| pos 0 | ✓ (1.0) | ✗ (-inf) | ✗ (-inf) | ✗ (-inf) | ✗ (-inf) |
| pos 1 | ✓ (1.0) | ✓ (1.0) | ✗ (-inf) | ✗ (-inf) | ✗ (-inf) |
| pos 2 | ✓ (1.0) | ✓ (1.0) | ✓ (1.0) | ✗ (-inf) | ✗ (-inf) |
| pos 3 | ✓ (1.0) | ✓ (1.0) | ✓ (1.0) | ✓ (1.0) | ✗ (-inf) |
| pos 4 | ✓ (1.0) | ✓ (1.0) | ✓ (1.0) | ✓ (1.0) | ✓ (1.0) |
The lower triangle (including diagonal) is unmasked. The strict upper triangle is -inf. PyTorch MHA's attn_mask argument accepts exactly this convention when is_causal=False.
Discrimination
Sort into buckets
Sort these by can see pos 1, from memory, without looking back at The causal mask matrix (T = 5). Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Ranking
Put in order
Put the moves of Masked self-attention: trace through T = 4 into the order they have to happen.
Why: These are the moves of the worked example in the order it makes them, and each one is set up by the one before it. Each of the four token embeddings gets mapped to a query, key, and value vector.
Worked example
Set up inputs: x shape (1, 4, 8), compute Q, K, V via linear projections
Why: Each of the four token embeddings gets mapped to a query, key, and value vector. Seed 42, d_model=8.
Compute raw scores: (Q @ K^T) / sqrt(8)
Why: Scale by sqrt(d_k)=2.828 to prevent vanishing softmax gradients (Lesson 83).
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(42)
T, d = 4, 8
x = torch.randn(1, T, d)
Wq = nn.Linear(d, d, bias=False)
Wk = nn.Linear(d, d, bias=False)
Wv = nn.Linear(d, d, bias=False)
Q, K, V = Wq(x), Wk(x), Wv(x)
scores = (Q @ K.transpose(-2,-1)) / d**0.5
print(scores[0].detach().numpy().round(3))| query | key 0 | key 1 | key 2 | key 3 |
|---|---|---|---|---|
| pos 0 | 0.443 | 0.015 | -0.145 | 0.213 |
| pos 1 | -0.234 | 0.039 | -0.067 | 0.246 |
| pos 2 | 0.122 | -0.212 | -0.201 | -0.866 |
| pos 3 | -0.104 | 0.343 | 0.656 | 0.011 |
Add causal mask: fill upper-triangle positions with -inf
Why: torch.triu(..., diagonal=1).bool() selects positions j > i. masked_fill replaces them with -inf.
| query | key 0 | key 1 | key 2 | key 3 |
|---|---|---|---|---|
| pos 0 | 0.443 | -inf | -inf | -inf |
| pos 1 | -0.234 | 0.039 | -inf | -inf |
| pos 2 | 0.122 | -0.212 | -0.201 | -inf |
| pos 3 | -0.104 | 0.343 | 0.656 | 0.011 |
Lines highlighted: build mask, apply fill.
Apply softmax row-wise: e^(-inf) = 0, so future positions vanish
Why: softmax normalises each row to sum=1. The -inf entries become exactly 0.0 — no gradient flows back through them.
| query | w to pos 0 | w to pos 1 | w to pos 2 | w to pos 3 | sum |
|---|---|---|---|---|---|
| pos 0 | 1.0000 | 0.0000 | 0.0000 | 0.0000 | 1.0 |
| pos 1 | 0.4323 | 0.5677 | 0.0000 | 0.0000 | 1.0 |
| pos 2 | 0.4099 | 0.2935 | 0.2966 | 0.0000 | 1.0 |
| pos 3 | 0.1717 | 0.2685 | 0.3672 | 0.1926 | 1.0 |
Comparison
Comparison matrix
From Masked self-attention: trace through T = 4: refill the key 1 column from what you know. The rest of the table is as it appeared.
| query | key 0 | key 1 | key 2 | key 3 |
|---|---|---|---|---|
| pos 0 | 0.443 | 0.015 | -0.145 | 0.213 |
| pos 1 | -0.234 | 0.039 | -0.067 | 0.246 |
| pos 2 | 0.122 | -0.212 | -0.201 | -0.866 |
| pos 3 | -0.104 | 0.343 | 0.656 | 0.011 |
Anomaly
Predict first
A student writes this, and it looks reasonable:
Compute softmax first, then zero out future positions in the attention weight matrix.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Seems to block future tokens — the upper triangle becomes 0.
Add -inf to the logits before softmax, so the exponential maps them to 0 naturally.
Why: Seems to block future tokens — the upper triangle becomes 0.
Trap
Compute softmax first, then zero out future positions in the attention weight matrix.
attn = softmax(scores); attn = attn * tril_mask
Why: Seems to block future tokens — the upper triangle becomes 0.
Row sums after zero-masking: [0.032, 0.119, 0.707, 1.000] — rows no longer sum to 1. The weighted-sum context vectors are scaled down arbitrarily, corrupting the representation.
Add -inf to the logits before softmax, so the exponential maps them to 0 naturally.
scores = scores.masked_fill(future_mask, float('-inf')); attn = softmax(scores)
Why: e^(-inf) = 0 exactly; softmax renormalises over only the visible positions; each row still sums to 1.
Row sums: [1.0, 1.0, 1.0, 1.0]. Attention weights remain a proper probability distribution over visible tokens.
Break the constraint
Discussion prompt
The rule this trap just fixed:
Add -inf to the logits before softmax, so the exponential maps them to 0 naturally.
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
Seems to block future tokens — the upper triangle becomes 0.
Section
Part 3 of 4
Concept
Self-attention uses a single sequence for Q, K, and V. Cross-attention uses two sequences: Q comes from the current decoder state; K and V come from the encoder output. The decoder token 'asks a question' (Q) and the encoder answers with its context (K/V).
\[ \text{CrossAttn}(x_{\text{dec}},\, z_{\text{enc}}) = \text{softmax}\!\left(\frac{Q_{\text{dec}}\,K_{\text{enc}}^\top}{\sqrt{d_k}}\right) V_{\text{enc}} \]
| matrix | source | shape |
|---|---|---|
| Q | decoder hidden state x_dec | (B, T_dec, d_k) |
| K | encoder output z_enc | (B, T_enc, d_k) |
| V | encoder output z_enc | (B, T_enc, d_v) |
| scores | Q @ K^T / sqrt(d_k) | (B, T_dec, T_enc) |
| output | softmax(scores) @ V | (B, T_dec, d_v) |
Trade off
Comparison matrix
From Cross-attention: asymmetric Q vs K/V: every row here is a choice with a cost. Fill the shape column, then say which row you would actually pick and what you give up for it.
| matrix | source | shape |
|---|---|---|
| Q | decoder hidden state x_dec | (B, T_dec, d_k) |
| K | encoder output z_enc | (B, T_enc, d_k) |
| V | encoder output z_enc | (B, T_enc, d_v) |
| scores | Q @ K^T / sqrt(d_k) | (B, T_dec, T_enc) |
| output | softmax(scores) @ V | (B, T_dec, d_v) |
Step zero
Discussion prompt
Cross-attention forward: T_dec=4, T_enc=6, d=8 — before any calculation: what is the plan? Name the moves in order, in plain English, without doing the arithmetic.
Hint: It starts with: Project decoder states and encoder output into Q, K, V spaces
Answer:
Worked example
Project decoder states and encoder output into Q, K, V spaces
Why: Q uses a decoder-side linear; K and V use encoder-side linears. All project to d=8.
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(42)
T_dec, T_enc, d = 4, 6, 8
dec = torch.randn(1, T_dec, d)
enc = torch.randn(1, T_enc, d)
Wq = nn.Linear(d, d, bias=False)
Wk = nn.Linear(d, d, bias=False)
Wv = nn.Linear(d, d, bias=False)
Q = Wq(dec) # (1,4,8)
K = Wk(enc) # (1,6,8)
V = Wv(enc) # (1,6,8)
scores = (Q @ K.transpose(-2,-1)) / d**0.5 # (1,4,6)
weights = F.softmax(scores, dim=-1)
out = weights @ V # (1,4,8)
print(weights[0].detach().numpy().round(3))Read off the score matrix shape: (B=1, T_dec=4, T_enc=6)
Why: Each decoder token (row) attends across all 6 encoder positions (columns). No causal mask needed here — the encoder output is fully visible.
| dec token | enc 0 | enc 1 | enc 2 | enc 3 | enc 4 | enc 5 |
|---|---|---|---|---|---|---|
| pos 0 | 0.122 | 0.393 | 0.134 | 0.115 | 0.129 | 0.106 |
| pos 1 | 0.206 | 0.165 | 0.155 | 0.122 | 0.212 | 0.140 |
| pos 2 | 0.160 | 0.258 | 0.167 | 0.150 | 0.143 | 0.123 |
| pos 3 | 0.149 | 0.141 | 0.170 | 0.156 | 0.168 | 0.216 |
Each row sums to 1.0; each dec token gets a blended encoder context vector of shape (8,)
Why: weights @ V is (1,4,6) @ (1,6,8) = (1,4,8). Output shape matches dec input shape — the residual connection can add them directly.
Error analysis
Annotate
Walk the callouts on Cross-attention forward: T_dec=4, T_enc=6, d=8. Each one is a place this is easy to get subtly wrong.
Section
Part 4 of 4
Concept
Each of the three sublayers is wrapped in a residual connection + LayerNorm (Lesson 84 pattern). The residual lets gradients bypass the sublayer entirely; LayerNorm stabilises activations across the token dimension.
\[ \begin{aligned} x_1 &= \text{LN}(x + \text{MaskedSelfAttn}(x)) \\ x_2 &= \text{LN}(x_1 + \text{CrossAttn}(x_1, z)) \\ x_3 &= \text{LN}(x_2 + \text{FFN}(x_2)) \end{aligned} \]
The output x_3 has the same shape as the input x — (B, T_dec, d_model) — so multiple decoder blocks can be stacked identically, just as encoder blocks are stacked (Lesson 84).
Counterexample
Discussion prompt
The output x_3 has the same shape as the input x — (B, T_dec, d_model) — so multiple decoder blocks can be stacked identically, just as encoder blocks are stacked (Lesson 84).
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Step zero
Discussion prompt
Implement and trace DecoderBlock: d=16, T_dec=4, T_enc=6 — before any calculation: what is the plan? Name the moves in order, in plain English, without doing the arithmetic.
Hint: It starts with: Define DecoderBlock with nn.MultiheadAttention for both attention…
Answer:
Worked example
Define DecoderBlock with nn.MultiheadAttention for both attention sublayers
Why: MHA handles the multi-head split internally. Setting batch_first=True keeps shapes (B, T, d) throughout.
import torch, torch.nn as nn
torch.manual_seed(42)
class DecoderBlock(nn.Module):
def __init__(self, d_model, n_heads, d_ff):
super().__init__()
self.masked_attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
self.cross_attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
self.ffn = nn.Sequential(
nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model))
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.norm3 = nn.LayerNorm(d_model)
def forward(self, x, encoder_out, causal_mask=None):
a, _ = self.masked_attn(x, x, x, attn_mask=causal_mask)
x = self.norm1(x + a) # sublayer 1
c, _ = self.cross_attn(x, encoder_out, encoder_out)
x = self.norm2(x + c) # sublayer 2
x = self.norm3(x + self.ffn(x)) # sublayer 3
return xBuild inputs and causal mask, then call forward; check output shape
Why: causal_mask is the upper-triangular bool matrix; output must be (1, 4, 16) to allow residual stacking.
| variable | shape | description |
|---|---|---|
| dec_input (x) | (1, 4, 16) | 4 decoder tokens, d_model=16 |
| enc_output (z) | (1, 6, 16) | 6 encoder positions |
| causal_mask | (4, 4) | upper-tri bool, blocks j > i |
| after sublayer 1 | (1, 4, 16) | post masked-self-attn + LN |
| after sublayer 2 | (1, 4, 16) | post cross-attn + LN |
| output x_3 | (1, 4, 16) | post FFN + LN; same as input |
Verify: dec_input[0,0] norm ≈ 4.250, output[0,0] norm ≈ 4.000
Why: LayerNorm contracts norms toward a normalised scale. The output norm being slightly smaller than the input confirms LN is operating correctly.
Comparison
Comparison matrix
From Implement and trace DecoderBlock: d=16, T_dec=4, T_enc=6: refill the shape column from what you know. The rest of the table is as it appeared.
| variable | shape | description |
|---|---|---|
| dec_input (x) | (1, 4, 16) | 4 decoder tokens, d_model=16 |
| enc_output (z) | (1, 6, 16) | 6 encoder positions |
| causal_mask | (4, 4) | upper-tri bool, blocks j > i |
| after sublayer 1 | (1, 4, 16) | post masked-self-attn + LN |
| after sublayer 2 | (1, 4, 16) | post cross-attn + LN |
| output x_3 | (1, 4, 16) | post FFN + LN; same as input |
Concept
At inference time there is no teacher forcing — the decoder has not yet seen future tokens. It generates one token per step: run the decoder on all tokens so far, take the last position's logits, argmax (or sample), append the new token, repeat.
\[ \hat{y}_t = \arg\max_v\, \text{head}\bigl(\text{DecoderBlock}(y_{<t},\, z)\bigr)_{t-1} \]
At step t the decoder takes t input tokens, runs the full block, and reads off the logit at position t-1 (the last position). The causal mask ensures position t-1 cannot attend to position t (which does not exist yet).
Sorting
Sort into buckets
These are the pieces of Lesson 85: Transformer Decoder Block, out of order. Put each one back under the part of the lesson it belongs to.
Ranking
Put in order
Put the moves of Autoregressive generation: step-by-step trace into the order they have to happen.
Why: These are the moves of the worked example in the order it makes them, and each one is set up by the one before it. The encoder runs only once. The decoder grows its context window token by token without re-encoding the source.
Worked example
Encode source once; start decoder with [BOS] token (id=1)
Why: The encoder runs only once. The decoder grows its context window token by token without re-encoding the source.
import torch, torch.nn as nn
torch.manual_seed(42)
class TinyTranslator(nn.Module):
def __init__(self, V, d, H, d_ff):
super().__init__()
self.embed = nn.Embedding(V, d)
self.block = DecoderBlock(d, H, d_ff)
self.head = nn.Linear(d, V)
def encode(self, src): return self.embed(src)
def decode_step(self, tgt, enc_out):
T = tgt.shape[1]
mask = torch.triu(torch.ones(T,T), diagonal=1).bool()
x = self.block(self.embed(tgt), enc_out, causal_mask=mask)
return self.head(x) # (1, T, V)
model = TinyTranslator(V=20, d=16, H=2, d_ff=32)
src = torch.tensor([[3,7,12,5,9]])
enc_out = model.encode(src) # (1,5,16)
generated = [1] # BOS=1
for step in range(4):
tgt = torch.tensor([generated])
logits = model.decode_step(tgt, enc_out)
tok = logits[0,-1].argmax().item()
generated.append(tok)
print(generated)Trace each step: input grows by one token; only the last logit is consumed
Why: The O(T^2) attention cost rises each step. KV-caching (future lesson) avoids recomputing past keys/values.
| step | tgt_ids fed in | tgt shape | logit read at pos |
|---|---|---|---|
| 1 | [1] | (1,1) | pos 0 |
| 2 | [1, tok1] | (1,2) | pos 1 |
| 3 | [1, t1, t2] | (1,3) | pos 2 |
| 4 | [1,t1,t2,t3] | (1,4) | pos 3 |
Stop when the predicted token equals EOS, or when max_length is reached
Why: There is no natural stopping criterion from the architecture itself — EOS is a special vocabulary token the model learns to predict during training.
Trade off
Comparison matrix
From Autoregressive generation: step-by-step trace: every row here is a choice with a cost. Fill the logit read at pos column, then say which row you would actually pick and what you give up for it.
| step | tgt_ids fed in | tgt shape | logit read at pos |
|---|---|---|---|
| 1 | [1] | (1,1) | pos 0 |
| 2 | [1, tok1] | (1,2) | pos 1 |
| 3 | [1, t1, t2] | (1,3) | pos 2 |
| 4 | [1,t1,t2,t3] | (1,4) | pos 3 |
Ranking
Put in order
These are the steps of The decoder block recipe, scrambled. Put them back in order before the next slide shows you.
torch.triu(ones(T,T), diagonal=1).bool() — shape (T, T)attn_mask to MHA; add residual; apply LN → x1Why: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.
Pattern
torch.triu(ones(T,T), diagonal=1).bool() — shape (T, T)attn_mask to MHA; add residual; apply LN → x1Edge cases
Discussion prompt
The decoder block recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
torch.triu(ones(T,T), diagonal=1).bool() — shape (T, T)attn_mask to MHA; add residual; apply LN → x1Elimination
Eliminate the wrong options
You apply a causal mask by zeroing out the upper-triangle of the attention weight matrix after softmax. Which failure mode does this introduce?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: After softmax, the probabilities over all T positions already sum to 1. Zeroing the upper triangle removes probability mass without renormalising — e.g. for pos 0 the only unmasked entry had weight 1.0 before zeroing, but if softmax ran on all 4 logits first, the single visible weight might be 0.032 (our real example). The context vector is then scaled by 0.032 instead of 1.0, corrupting the representation.
Check
Reason through this before clicking.
Check your understanding
You apply a causal mask by zeroing out the upper-triangle of the attention weight matrix after softmax. Which failure mode does this introduce?
Answer: A
Why: After softmax, the probabilities over all T positions already sum to 1. Zeroing the upper triangle removes probability mass without renormalising — e.g. for pos 0 the only unmasked entry had weight 1.0 before zeroing, but if softmax ran on all 4 logits first, the single visible weight might be 0.032 (our real example). The context vector is then scaled by 0.032 instead of 1.0, corrupting the representation.
Prediction
Predict first
In cross-attention with T_dec=4, T_enc=6, d_model=8, n_heads=2: what is the shape of the attention weight matrix (after softmax, before multiplying V)?
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: (B, 2, 4, 6) — one head per row in each head's subspace
Why: With n_heads=2 and d_model=8 each head has d_k=4. Within each head, Q_dec is (B, 2, 4, 4) and K_enc is (B, 2, 6, 4), so Q@K^T is (B, 2, 4, 6). PyTorch's MHA stacks these and returns the merged output, but internally the score/weight tensor is (B, n_heads, T_dec, T_enc) = (B, 2, 4, 6).
Check
Trace the shapes carefully.
Check your understanding
In cross-attention with T_dec=4, T_enc=6, d_model=8, n_heads=2: what is the shape of the attention weight matrix (after softmax, before multiplying V)?
Answer: A
Why: With n_heads=2 and d_model=8 each head has d_k=4. Within each head, Q_dec is (B, 2, 4, 4) and K_enc is (B, 2, 6, 4), so Q@K^T is (B, 2, 4, 6). PyTorch's MHA stacks these and returns the merged output, but internally the score/weight tensor is (B, n_heads, T_dec, T_enc) = (B, 2, 4, 6).
Elimination
Eliminate the wrong options
During autoregressive decoding, you have generated tokens [BOS, 7, 3] so far. The encoder output is fixed at shape (1, 5, d). What does the decoder receive as input for the next step?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: B
Why: Standard autoregressive decoding re-feeds the entire prefix at each step. The decoder runs on [BOS, 7, 3] (shape (1,3)), producing three output positions; we only use the logit at position 2 (the last) to predict token 4. The causal mask is 3×3 (lower-triangular). KV-caching avoids recomputing earlier keys/values, but the logical computation is identical.
Check
Consider what the decoder actually runs at each step.
Check your understanding
During autoregressive decoding, you have generated tokens [BOS, 7, 3] so far. The encoder output is fixed at shape (1, 5, d). What does the decoder receive as input for the next step?
Answer: B
Why: Standard autoregressive decoding re-feeds the entire prefix at each step. The decoder runs on [BOS, 7, 3] (shape (1,3)), producing three output positions; we only use the logit at position 2 (the last) to predict token 4. The causal mask is 3×3 (lower-triangular). KV-caching avoids recomputing earlier keys/values, but the logical computation is identical.
Section
Project
Concept
You will implement a complete seq2seq decoder stack and verify it generates coherent token sequences on a toy vocabulary. No pretrained weights — build everything from scratch using random embeddings and verify shapes at every step.
build_causal_mask(T) and verify the 5×5 output matches the expected patternDecoderBlock (masked-self-attn + cross-attn + FFN + LayerNorm) and check output shapeautoregressive_decode(model, src, max_len) that generates a token sequence and stops at EOSStep zero
Discussion prompt
Milestone A: build_causal_mask — before any calculation: what is the plan? Name the moves in order, in plain English, without doing the arithmetic.
Hint: It starts with: Task: write build_causal_mask(T) returning upper-triangular bool mask
Answer:
Worked example
Task: write build_causal_mask(T) returning upper-triangular bool mask
Why: Use torch.triu with diagonal=1 to select positions j > i. dtype must be bool for MHA's attn_mask.
Hint: torch.triu(torch.ones(T, T), diagonal=1).bool()
Why: diagonal=1 means the main diagonal (i==j) is NOT masked — each position can attend to itself.
Solution and expected output
Why: For T=5 the mask is 5×5; row 0 has 4 True entries (blocks positions 1-4), row 4 has 0 True entries.
import torch
def build_causal_mask(T):
return torch.triu(torch.ones(T, T), diagonal=1).bool()
mask = build_causal_mask(5)
print(mask.int())| row (query) | col 0 | col 1 | col 2 | col 3 | col 4 |
|---|---|---|---|---|---|
| 0 | 0 | 1 | 1 | 1 | 1 |
| 1 | 0 | 0 | 1 | 1 | 1 |
| 2 | 0 | 0 | 0 | 1 | 1 |
| 3 | 0 | 0 | 0 | 0 | 1 |
| 4 | 0 | 0 | 0 | 0 | 0 |
Error analysis
Annotate
Walk the callouts on Milestone A: build_causal_mask. Each one is a place this is easy to get subtly wrong.
Ranking
Put in order
Put the moves of Milestone B: DecoderBlock shape check into the order they have to happen.
Why: These are the moves of the worked example in the order it makes them, and each one is set up by the one before it. If the shapes match you can stack N blocks.
Worked example
Task: implement DecoderBlock and verify output shape equals input shape
Why: If the shapes match you can stack N blocks. Use d_model=16, n_heads=2, d_ff=32.
Hint: pass causal_mask as attn_mask to masked_attn; no mask for cross_attn
Why: The cross-attention query can attend to all encoder positions — the encoder output is fully available at decode time.
Solution: run forward and print shapes at each sublayer
Why: All three intermediate tensors and the final output share shape (1, 4, 16), confirming each residual connection is valid.
import torch, torch.nn as nn
torch.manual_seed(42)
class DecoderBlock(nn.Module):
def __init__(self, d, H, d_ff):
super().__init__()
self.m_attn = nn.MultiheadAttention(d, H, batch_first=True)
self.c_attn = nn.MultiheadAttention(d, H, batch_first=True)
self.ffn = nn.Sequential(nn.Linear(d,d_ff),nn.ReLU(),nn.Linear(d_ff,d))
self.ln1, self.ln2, self.ln3 = (nn.LayerNorm(d) for _ in range(3))
def forward(self, x, z, mask=None):
a,_ = self.m_attn(x, x, x, attn_mask=mask)
x = self.ln1(x + a)
c,_ = self.c_attn(x, z, z)
x = self.ln2(x + c)
return self.ln3(x + self.ffn(x))
blk = DecoderBlock(16, 2, 32)
x = torch.randn(1,4,16)
z = torch.randn(1,6,16)
msk = torch.triu(torch.ones(4,4),diagonal=1).bool()
out = blk(x, z, mask=msk)
print('in:', x.shape, ' out:', out.shape)| tensor | shape | note |
|---|---|---|
| x (input) | (1, 4, 16) | dec tokens |
| z (encoder_out) | (1, 6, 16) | encoder tokens |
| mask | (4, 4) | upper-tri bool |
| out (output) | (1, 4, 16) | same as x — stackable |
Discrimination
Sort into buckets
Sort these by shape, from memory, without looking back at Milestone B: DecoderBlock shape check. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Pattern
Predict first
The table runs: 1 | 1 ([BOS]) | (1,1,20) | argmax of pos 0 · 2 | 2 | (1,2,20) | argmax of pos 1 · 3 | 3 | (1,3,20) | argmax of pos 2 · 4 | 4 | (1,4,20) | argmax of pos 3
In Milestone C + Show-it-off: autoregressive decode & stacked…, given the rows so far: what is the next one — the row where step is 5?
Correct: 5 | 5 | (1,5,20) | argmax of pos 4
| step | tgt len | logit shape | next token |
|---|---|---|---|
| 1 | 1 ([BOS]) | (1,1,20) | argmax of pos 0 |
| 2 | 2 | (1,2,20) | argmax of pos 1 |
| 3 | 3 | (1,3,20) | argmax of pos 2 |
| 4 | 4 | (1,4,20) | argmax of pos 3 |
| 5 | 5 | (1,5,20) | argmax of pos 4 |
Why: The relationship between the columns, not the individual numbers, is what generates the next row. Start with [BOS]; at each step run decode_step on the full prefix; append argmax of the last logit; stop at EOS or max_len.
Worked example
Task: implement autoregressive_decode(model, src, max_len=8) returning a token list
Why: Start with [BOS]; at each step run decode_step on the full prefix; append argmax of the last logit; stop at EOS or max_len.
Show-it-off: wrap two DecoderBlocks into a stack; print output shape matches input shape
Why: Stacking is the standard transformer pattern. GPT-2 small uses 12 decoder blocks; GPT-4 architecture reportedly uses ~96.
import torch, torch.nn as nn
torch.manual_seed(42)
class DecoderStack(nn.Module):
def __init__(self, V, d, H, d_ff, n_layers):
super().__init__()
self.embed = nn.Embedding(V, d)
self.blocks = nn.ModuleList([DecoderBlock(d,H,d_ff) for _ in range(n_layers)])
self.head = nn.Linear(d, V)
def forward(self, tgt, enc_out):
T = tgt.shape[1]
msk = torch.triu(torch.ones(T,T),diagonal=1).bool()
x = self.embed(tgt)
for blk in self.blocks:
x = blk(x, enc_out, mask=msk)
return self.head(x)
V = 20; d = 16; enc_out = torch.randn(1,5,d)
model = DecoderStack(V, d, H=2, d_ff=32, n_layers=2)
generated = [1] # BOS
for _ in range(5):
tgt = torch.tensor([generated])
logits = model(tgt, enc_out)
nxt = logits[0,-1].argmax().item()
generated.append(nxt)
if nxt == 2: break # EOS
print('Generated:', generated)| step | tgt len | logit shape | next token |
|---|---|---|---|
| 1 | 1 ([BOS]) | (1,1,20) | argmax of pos 0 |
| 2 | 2 | (1,2,20) | argmax of pos 1 |
| 3 | 3 | (1,3,20) | argmax of pos 2 |
| 4 | 4 | (1,4,20) | argmax of pos 3 |
| 5 | 5 | (1,5,20) | argmax of pos 4 |
Comparison
Comparison matrix
From Milestone C + Show-it-off: autoregressive decode & stacked…: refill the logit shape column from what you know. The rest of the table is as it appeared.
| step | tgt len | logit shape | next token |
|---|---|---|---|
| 1 | 1 ([BOS]) | (1,1,20) | argmax of pos 0 |
| 2 | 2 | (1,2,20) | argmax of pos 1 |
| 3 | 3 | (1,3,20) | argmax of pos 2 |
| 4 | 4 | (1,4,20) | argmax of pos 3 |
| 5 | 5 | (1,5,20) | argmax of pos 4 |
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — The three sublayers · Causal masked self-attention · Cross-attention: decoder reads the encoder · DecoderBlock: full forward pass · Your turn: build a DecoderBlock. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
| concept | key formula / shape | lesson callback |
|---|---|---|
| causal mask | triu(ones(T,T), diag=1).bool() | Lesson 83 (attention) |
| masked self-attn | Q=K=V=x; scores shape (T,T) | Lesson 83 |
| cross-attn scores | (B, T_dec, T_enc) | Lesson 83 (dot-product attn) |
| softmax on -inf | exp(-inf)=0 → row still sums to 1 | Lesson 17 (softmax+CE) |
| FFN sublayer | Linear→ReLU→Linear per position | Lesson 84 (encoder) |
| autoregressive loop | grow prefix; argmax last logit | Lesson 85 (this lesson) |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.