Lesson 85: Transformer Decoder Block

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

What this lesson covers

The lesson, slide by slide

1. Transformer Decoder Block

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.

2. By the end of this lesson you can

Objectives

  1. Name the three sublayers of a decoder block and state their order
  2. Build the causal (upper-triangular) mask and explain why -inf goes in the logits, not the outputs
  3. Implement cross-attention with Q from the decoder and K/V from the encoder
  4. Trace a full DecoderBlock forward pass through shapes for arbitrary T_dec, T_enc, d_model
  5. Describe autoregressive generation: one token per step, growing context window

3. What survived from Transformer Encoder for Classification?

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.

4. The three sublayers

Section

Part 1 of 4

5. Decoder block anatomy

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).

① Masked self-attention
Attends only to past/current tokens. Causal mask blocks future peeking.
② Cross-attention
Q comes from the decoder; K and V come from the encoder output. Decoder reads the source.
③ FFN
Same position-wise two-layer MLP as in the encoder (Lesson 84). Expands to d_ff then contracts.

\[ \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) \]

6. Which is which: Decoder block anatomy

Matching

Match the pairs

From Decoder block anatomy — match each one to what it actually does. The descriptions have been shuffled.

  • c1. ① Masked self-attention
  • c2. ② Cross-attention
  • c3. ③ FFN
  • b1. Attends only to past/current tokens. Causal mask blocks future peeking.
  • b2. Q comes from the decoder; K and V come from the encoder output. Decoder reads the source.
  • b3. Same position-wise two-layer MLP as in the encoder (Lesson 84). Expands to d_ff then contracts.

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.

7. Causal masked self-attention

Section

Part 2 of 4

8. Why mask? Preventing future peek

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.

9. Break it if you can: Why mask? Preventing future peek

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.

10. The causal mask matrix (T = 5)

Concept

tokencan see pos 0can see pos 1can see pos 2can see pos 3can 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.

11. Which is which, by can see pos 1

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.

✗ (-inf)
pos 0
✓ (1.0)
pos 1; pos 2; pos 3; pos 4
g1
can see pos 1 is "✗ (-inf)" for pos 0 — that is what the table on "The causal mask matrix (T = 5)" records, and it is the single property separating this group from the rest.
g2
can see pos 1 is "✓ (1.0)" for pos 1, pos 2, pos 3, pos 4 — that is what the table on "The causal mask matrix (T = 5)" records, and it is the single property separating this group from the rest.

12. What has to happen first: Masked self-attention: trace through T = 4

Ranking

Put in order

Put the moves of Masked self-attention: trace through T = 4 into the order they have to happen.

  1. Set up inputs: x shape (1, 4, 8), compute Q, K, V via linear projections
  2. Compute raw scores: (Q @ K^T) / sqrt(8)
  3. Add causal mask: fill upper-triangle positions with -inf
  4. Apply softmax row-wise: e^(-inf) = 0, so future positions vanish

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.

13. Masked self-attention: trace through T = 4

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))
querykey 0key 1key 2key 3
pos 00.4430.015-0.1450.213
pos 1-0.2340.039-0.0670.246
pos 20.122-0.212-0.201-0.866
pos 3-0.1040.3430.6560.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.

querykey 0key 1key 2key 3
pos 00.443-inf-inf-inf
pos 1-0.2340.039-inf-inf
pos 20.122-0.212-0.201-inf
pos 3-0.1040.3430.6560.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.

queryw to pos 0w to pos 1w to pos 2w to pos 3sum
pos 01.00000.00000.00000.00001.0
pos 10.43230.56770.00000.00001.0
pos 20.40990.29350.29660.00001.0
pos 30.17170.26850.36720.19261.0

14. Fill in: key 1 for Masked self-attention: trace through T = 4

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.

querykey 0key 1key 2key 3
pos 00.4430.015-0.1450.213
pos 1-0.2340.039-0.0670.246
pos 20.122-0.212-0.201-0.866
pos 3-0.1040.3430.6560.011

15. Something is wrong here: masking the attention outputs instead of the logits

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.

16. Trap: masking the attention outputs instead of the logits

Trap

The 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.

The fix

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.

17. Break it on purpose: masking the attention outputs instead of the…

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.

18. Cross-attention: decoder reads the encoder

Section

Part 3 of 4

19. Cross-attention: asymmetric Q vs K/V

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}} \]

matrixsourceshape
Qdecoder hidden state x_dec(B, T_dec, d_k)
Kencoder output z_enc(B, T_enc, d_k)
Vencoder output z_enc(B, T_enc, d_v)
scoresQ @ K^T / sqrt(d_k)(B, T_dec, T_enc)
outputsoftmax(scores) @ V(B, T_dec, d_v)

20. What each one costs: Cross-attention: asymmetric Q vs K/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.

matrixsourceshape
Qdecoder hidden state x_dec(B, T_dec, d_k)
Kencoder output z_enc(B, T_enc, d_k)
Vencoder output z_enc(B, T_enc, d_v)
scoresQ @ K^T / sqrt(d_k)(B, T_dec, T_enc)
outputsoftmax(scores) @ V(B, T_dec, d_v)

21. Plan first: Cross-attention forward: T_dec=4, T_enc=6, d=8

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:

  1. Project decoder states and encoder output into Q, K, V spaces
  2. Read off the score matrix shape: (B=1, T_dec=4, T_enc=6)
  3. Each row sums to 1.0; each dec token gets a blended encoder context vector of shape (8,)

22. Cross-attention forward: T_dec=4, T_enc=6, d=8

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 tokenenc 0enc 1enc 2enc 3enc 4enc 5
pos 00.1220.3930.1340.1150.1290.106
pos 10.2060.1650.1550.1220.2120.140
pos 20.1600.2580.1670.1500.1430.123
pos 30.1490.1410.1700.1560.1680.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.

23. Inspect it line by line: Cross-attention forward: T_dec=4, T_enc=6…

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.

  • Q uses a decoder-side linear; K and V use encoder-side linears. All project to d=8.
  • Each decoder token (row) attends across all 6 encoder positions (columns). No causal mask needed here — the encoder output is fully visible.
  • 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.

24. DecoderBlock: full forward pass

Section

Part 4 of 4

25. DecoderBlock: residuals and LayerNorm

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).

26. Break it if you can: DecoderBlock: residuals and LayerNorm

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.

27. Plan first: Implement and trace DecoderBlock: d=16, T_dec=4, T_enc=6

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:

  1. Define DecoderBlock with nn.MultiheadAttention for both attention sublayers
  2. Build inputs and causal mask, then call forward; check output shape
  3. Verify: dec_input[0,0] norm ≈ 4.250, output[0,0] norm ≈ 4.000

28. Implement and trace DecoderBlock: d=16, T_dec=4, T_enc=6

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 x

Build 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.

variableshapedescription
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.

29. Fill in: shape for Implement and trace DecoderBlock: d=16…

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.

variableshapedescription
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

30. Autoregressive generation

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).

31. Where does each piece belong: Lesson 85: Transformer Decoder Block

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.

Causal masked self-attention
Why mask? Preventing future peek; The causal mask matrix (T = 5); Masked self-attention: trace through T = 4
Cross-attention: decoder reads the encoder
Cross-attention: asymmetric Q vs K/V; Cross-attention forward: T_dec=4, T_enc=6, d=8
DecoderBlock: full forward pass
DecoderBlock: residuals and LayerNorm; Implement and trace DecoderBlock: d=16, T_dec=4, T_enc=6; Autoregressive generation
s1
Causal masked self-attention is where Lesson 85: Transformer Decoder Block puts Why mask? Preventing future peek, The causal mask matrix (T = 5), Masked self-attention: trace through T = 4. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
Cross-attention: decoder reads the encoder is where Lesson 85: Transformer Decoder Block puts Cross-attention: asymmetric Q vs K/V, Cross-attention forward: T_dec=4, T_enc=6, d=8. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
DecoderBlock: full forward pass is where Lesson 85: Transformer Decoder Block puts DecoderBlock: residuals and LayerNorm, Implement and trace DecoderBlock: d=16, T_dec=4, T_enc=6, Autoregressive generation. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

32. What has to happen first: Autoregressive generation: step-by-step trace

Ranking

Put in order

Put the moves of Autoregressive generation: step-by-step trace into the order they have to happen.

  1. Encode source once; start decoder with [BOS] token (id=1)
  2. Trace each step: input grows by one token; only the last logit is consumed
  3. Stop when the predicted token equals EOS, or when max_length is reached

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.

33. Autoregressive generation: step-by-step trace

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.

steptgt_ids fed intgt shapelogit 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.

34. What each one costs: Autoregressive generation: step-by-step trace

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.

steptgt_ids fed intgt shapelogit 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

35. Rebuild the recipe: The decoder block recipe

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.

  1. Build the causal mask: torch.triu(ones(T,T), diagonal=1).bool() — shape (T, T)
  2. Masked self-attention: Q=K=V=x; pass attn_mask to MHA; add residual; apply LN → x1
  3. Cross-attention: Q=x1, K=V=encoder_out (no causal mask); add residual; apply LN → x2
  4. FFN: Linear(d→d_ff)→ReLU→Linear(d_ff→d); add residual; apply LN → x3
  5. Output x3 has same shape as input x: stack N blocks identically
  6. Autoregressive decode: re-run decoder on growing prefix; read logit at last position

Why: 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.

36. The decoder block recipe

Pattern

  1. Build the causal mask: torch.triu(ones(T,T), diagonal=1).bool() — shape (T, T)
  2. Masked self-attention: Q=K=V=x; pass attn_mask to MHA; add residual; apply LN → x1
  3. Cross-attention: Q=x1, K=V=encoder_out (no causal mask); add residual; apply LN → x2
  4. FFN: Linear(d→d_ff)→ReLU→Linear(d_ff→d); add residual; apply LN → x3
  5. Output x3 has same shape as input x: stack N blocks identically
  6. Autoregressive decode: re-run decoder on growing prefix; read logit at last position

37. Where does it stop working: The decoder block recipe

Edge 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:

  1. Build the causal mask: torch.triu(ones(T,T), diagonal=1).bool() — shape (T, T)
  2. Masked self-attention: Q=K=V=x; pass attn_mask to MHA; add residual; apply LN → x1
  3. Cross-attention: Q=x1, K=V=encoder_out (no causal mask); add residual; apply LN → x2
  4. FFN: Linear(d→d_ff)→ReLU→Linear(d_ff→d); add residual; apply LN → x3
  5. Output x3 has same shape as input x: stack N blocks identically
  6. Autoregressive decode: re-run decoder on growing prefix; read logit at last position

38. Rule out three: Check 1: causal mask placement

Elimination

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.

  • A. Each row no longer sums to 1; the context vector is scaled incorrectly
  • B. Gradient flow to the visible positions is blocked
  • C. The causal mask is applied twice because softmax already zeroes future entries
  • D. The upper triangle gets non-zero weight because softmax maps -inf to 0.5

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.

39. Check 1: causal mask placement

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?

  • A. Each row no longer sums to 1; the context vector is scaled incorrectly (correct)
  • B. Gradient flow to the visible positions is blocked
  • C. The causal mask is applied twice because softmax already zeroes future entries
  • D. The upper triangle gets non-zero weight because softmax maps -inf to 0.5

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.

Why B tempts people
Gradients flow through the multiplication by the binary mask just fine — zeroing a weight doesn't block backprop through the remaining weights. The problem is the broken normalisation, not gradient flow.
Why C tempts people
Softmax never zeroes anything. It maps every finite input to a positive probability; only an input of -inf maps to 0. If no -inf was added to the logits, future positions still get positive weight.
Why D tempts people
softmax(-inf) = exp(-inf)/Z = 0/Z = 0, not 0.5. Softmax can only output 0.5 for an entry if the input equals the log-mean of the other entries; -inf always maps to exactly 0.

40. Answer it before you see the options: Check 2: cross-attention shapes

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).

41. Check 2: cross-attention shapes

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)?

  • A. (B, 2, 4, 6) — one head per row in each head's subspace (correct)
  • B. (B, 4, 4) — square because T_dec decoder tokens attend to each other
  • C. (B, 6, 6) — square because keys and values both come from the encoder
  • D. (B, 4, 8) — rows are decoder tokens, columns are d_model dimensions

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).

Why B tempts people
T_dec × T_dec would be the shape for self-attention (where Q and K come from the same sequence). Cross-attention queries the encoder, so the columns span T_enc encoder positions, not T_dec.
Why C tempts people
Both K and V come from the encoder, but the rows of the score matrix are queries — and queries come from the decoder. Row count = T_dec = 4, not T_enc = 6.
Why D tempts people
d_model=8 is the embedding dimension, not the sequence dimension. After softmax the weight matrix operates on sequence positions (T_enc=6), not embedding dimensions.

42. Rule out three: Check 3: autoregressive generation

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.

  • A. The single token 3 (the last generated token) — shape (1,1,d)
  • B. All tokens so far [BOS, 7, 3] — shape (1,3,d), with a 3×3 causal mask
  • C. The encoder output directly; the decoder just projects it to logits
  • D. Tokens [7, 3] only — BOS is discarded after the first step to save compute

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.

43. Check 3: autoregressive generation

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?

  • A. The single token 3 (the last generated token) — shape (1,1,d)
  • B. All tokens so far [BOS, 7, 3] — shape (1,3,d), with a 3×3 causal mask (correct)
  • C. The encoder output directly; the decoder just projects it to logits
  • D. Tokens [7, 3] only — BOS is discarded after the first step to save compute

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.

Why A tempts people
Feeding only the last token would discard all positional context built up in earlier self-attention passes. Without the full prefix the causal mask would always be a 1×1 matrix, and the decoder could not model dependencies on earlier generated tokens.
Why C tempts people
The encoder output is the source of K and V for cross-attention, not the decoder's input sequence. The decoder's input is the partial target sequence (the tokens generated so far).
Why D tempts people
BOS is retained throughout decoding. Discarding it would shift token positions and break the positional encoding alignment. Some architectures use a different start token convention, but the full prefix is always kept.

44. Your turn: build a DecoderBlock

Section

Project

45. Project brief: seq2seq decoder from scratch

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.

  1. Milestone A: implement build_causal_mask(T) and verify the 5×5 output matches the expected pattern
  2. Milestone B: implement DecoderBlock (masked-self-attn + cross-attn + FFN + LayerNorm) and check output shape
  3. Milestone C: implement autoregressive_decode(model, src, max_len) that generates a token sequence and stops at EOS
  4. Show it off: stack two decoder blocks; confirm that the output of block 2 has the same shape as the input to block 1

46. Plan first: Milestone A: build_causal_mask

Step 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:

  1. Task: write build_causal_mask(T) returning upper-triangular bool mask
  2. Hint: torch.triu(torch.ones(T, T), diagonal=1).bool()
  3. Solution and expected output

47. Milestone A: build_causal_mask

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 0col 1col 2col 3col 4
001111
100111
200011
300001
400000

48. Inspect it line by line: Milestone A: build_causal_mask

Error analysis

Annotate

Walk the callouts on Milestone A: build_causal_mask. Each one is a place this is easy to get subtly wrong.

  • Use torch.triu with diagonal=1 to select positions j > i. dtype must be bool for MHA's attn_mask.
  • diagonal=1 means the main diagonal (i==j) is NOT masked — each position can attend to itself.
  • For T=5 the mask is 5×5; row 0 has 4 True entries (blocks positions 1-4), row 4 has 0 True entries.

49. What has to happen first: Milestone B: DecoderBlock shape check

Ranking

Put in order

Put the moves of Milestone B: DecoderBlock shape check into the order they have to happen.

  1. Task: implement DecoderBlock and verify output shape equals input shape
  2. Hint: pass causal_mask as attn_mask to masked_attn; no mask for cross_attn
  3. Solution: run forward and print shapes at each sublayer

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.

50. Milestone B: DecoderBlock shape check

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)
tensorshapenote
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

51. Which is which, by shape

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.

(1, 4, 16)
x (input); out (output)
(1, 6, 16)
z (encoder_out)
(4, 4)
mask
g1
shape is "(1, 4, 16)" for x (input), out (output) — that is what the table on "Milestone B: DecoderBlock shape check" records, and it is the single property separating this group from the rest.
g2
shape is "(1, 6, 16)" for z (encoder_out) — that is what the table on "Milestone B: DecoderBlock shape check" records, and it is the single property separating this group from the rest.
g3
shape is "(4, 4)" for mask — that is what the table on "Milestone B: DecoderBlock shape check" records, and it is the single property separating this group from the rest.

52. Predict the next row: Milestone C + Show-it-off: autoregressive decode…

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

steptgt lenlogit shapenext token
11 ([BOS])(1,1,20)argmax of pos 0
22(1,2,20)argmax of pos 1
33(1,3,20)argmax of pos 2
44(1,4,20)argmax of pos 3
55(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.

53. Milestone C + Show-it-off: autoregressive decode & stacked blocks

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)
steptgt lenlogit shapenext token
11 ([BOS])(1,1,20)argmax of pos 0
22(1,2,20)argmax of pos 1
33(1,3,20)argmax of pos 2
44(1,4,20)argmax of pos 3
55(1,5,20)argmax of pos 4

54. Fill in: logit shape for Milestone C + Show-it-off: autoregressive…

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.

steptgt lenlogit shapenext token
11 ([BOS])(1,1,20)argmax of pos 0
22(1,2,20)argmax of pos 1
33(1,3,20)argmax of pos 2
44(1,4,20)argmax of pos 3
55(1,5,20)argmax of pos 4

55. Connect it up: Lesson 85: Transformer Decoder Block

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.

56. Lesson 85 recap

Recap

conceptkey formula / shapelesson callback
causal masktriu(ones(T,T), diag=1).bool()Lesson 83 (attention)
masked self-attnQ=K=V=x; scores shape (T,T)Lesson 83
cross-attn scores(B, T_dec, T_enc)Lesson 83 (dot-product attn)
softmax on -infexp(-inf)=0 → row still sums to 1Lesson 17 (softmax+CE)
FFN sublayerLinear→ReLU→Linear per positionLesson 84 (encoder)
autoregressive loopgrow prefix; argmax last logitLesson 85 (this lesson)

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 85 — Transformer Decoder Block — Barron · USAAIO Round 2 Preparation, 2026
  2. Causal mask, masked self-attention weights, cross-attention shapes, DecoderBlock forward, and autoregressive generation verified with torch 2.7.1+cpu, numpy 2.2.6, seed 42, June 2026 — Real execution, verified
  3. Vaswani et al., Attention Is All You Need, NeurIPS 2017 — arXiv:1706.03762

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108