Lesson 104: Text Decoding Strategies

USAAIO Lesson 104, from Phase 3. It covers greedy decoding, beam search with width-k expansion, length normalization, temperature scaling, top-k masking, and top-p, or nucleus, sampling, all verified with real PyTorch log-softmax arithmetic. It proves by counter-example that greedy decoding is suboptimal, and builds a full beam-search trace table from scratch for k=2 over two steps. The lesson runs to 30 slides.

Subject: Machine Learning · 63 slides · code lesson

Open the interactive version of this deck · Homework for this lesson

What this lesson covers

The lesson, slide by slide

1. Text Decoding Strategies

Title

USAAIO · Lesson 104 · Phase 3

Greedy, beam search, length normalization, temperature, top-k, and top-p — the full toolkit for turning a language model's logits into text. Every probability in this deck came from running real torch log-softmax arithmetic.

2. By the end of this lesson you can

Objectives

  1. Implement greedy decoding and prove by counter-example that it is not guaranteed to find the highest-probability sequence
  2. Execute beam search (width k) step-by-step: expand each beam, score all V continuations, keep top-k — with concrete log-prob arithmetic
  3. Apply length normalization (Google NMT alpha) and explain why raw log-prob penalizes long sequences
  4. Sample with temperature T, reason about the entropy trade-off (T→0 is greedy, T→∞ is uniform), and implement top-k and top-p (nucleus) masking
  5. Choose the right decoder for a given task (translation vs creative generation vs constrained output)

3. What survived from Seq2Seq Encoder-Decoder with Bahdanau Attention?

Warm-up

Discussion prompt

Before we open Lesson 104: Text Decoding Strategies: without looking back, what was the main idea of Seq2Seq Encoder-Decoder with Bahdanau Attention, 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:

seq2seq encoder-decoder architecture (GRU encoder reads source, GRU decoder generates target one token at a time), Bahdanau additive attention (energy scores, softmax alpha, context vector), teacher forcing for fast training convergence, exposure bias and the train-inference mismatch, scheduled sampling to bridge the gap, and the copy mechanism as an attention-based pointer to source tokens.

4. Greedy decoding — fast but fragile

Section

Part 1 of 4

5. Greedy decoding: argmax at every step

Concept

At each time step t the model outputs a logit vector of length V (vocab size). Greedy decoding picks the single token with the highest probability and feeds it back as input for step t+1.

\[ \hat{w}_t = \arg\max_{v} P(w_t = v \mid w_1, \ldots, w_{t-1}) \]

The full sequence probability is the product of per-step greedy probabilities — one choice per step, no look-ahead. Speed is O(V · T) — one softmax per step.

6. Break it if you can: Greedy decoding: argmax at every step

Counterexample

Discussion prompt

At each time step t the model outputs a logit vector of length V (vocab size). Greedy decoding picks the single token with the highest probability and feeds it back as input for step t+1.

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.

Answer:

The full sequence probability is the product of per-step greedy probabilities — one choice per step, no look-ahead. Speed is O(V · T) — one softmax per step.

7. Guess the shape of the answer: Greedy trace — two-step example

Estimation

Predict first

Vocab = [cat, sat, mat, bat, hat, fat]. Logits are real torch tensors; probs from F.softmax.

Commit before you compute: what does Greedy trace — two-step example come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: Output: 'cat sat', joint probability 0.3322

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. P(cat) = 0.4680, P(sat|cat) = 0.7098.

8. Greedy trace — two-step example

Worked example

Vocab = [cat, sat, mat, bat, hat, fat]. Logits are real torch tensors; probs from F.softmax.

import torch, torch.nn.functional as F
vocab = ["cat","sat","mat","bat","hat","fat"]
logits0 = torch.tensor([2.1, 1.5, 0.8, 0.3,-0.2,-0.9])
logits1 = torch.tensor([0.5, 3.2, 0.9,-0.1, 0.7, 1.1])
p0 = F.softmax(logits0, dim=-1)
p1 = F.softmax(logits1, dim=-1)
greedy0 = p0.argmax().item()   # 0 -> 'cat'
greedy1 = p1.argmax().item()   # 1 -> 'sat'
seq_prob = (p0[greedy0] * p1[greedy1]).item()
print(vocab[greedy0], vocab[greedy1], seq_prob)
stepgreedy tokenP(token)cumulative P
0cat0.46800.4680
1sat0.70980.3322
—sequence—0.3322

Output: 'cat sat', joint probability 0.3322

Why: P(cat) = 0.4680, P(sat|cat) = 0.7098. Joint = 0.4680 × 0.7098 = 0.3322. This is the product of two argmax choices — no backtracking ever occurs.

9. Fill in: greedy token for Greedy trace — two-step example

Comparison

Comparison matrix

From Greedy trace — two-step example: refill the greedy token column from what you know. The rest of the table is as it appeared.

stepgreedy tokenP(token)cumulative P
0cat0.46800.4680
1sat0.70980.3322
—sequence—0.3322

10. Why greedy is not optimal — the cliff argument

Concept

Greedy maximizes each local step. But the highest-probability sequence may start with a lower-probability token that opens a high-probability continuation. This is the classic cliff: a big step-0 probability drops off a cliff at step 1.

pathP(tok 0)P(tok 1 | tok 0)joint P
A → D (greedy)0.450.400.1800
B → D (optimal)0.400.850.3400

Greedy picks A at step 0 (P=0.45 > 0.40). But B leads to a continuation with P=0.85 vs 0.40, giving joint 0.34 > 0.18. Greedy loses by a factor of 1.9× on this two-token example.

11. What each one costs: Why greedy is not optimal — the cliff argument

Trade off

Comparison matrix

From Why greedy is not optimal — the cliff argument: every row here is a choice with a cost. Fill the joint P column, then say which row you would actually pick and what you give up for it.

pathP(tok 0)P(tok 1 | tok 0)joint P
A → D (greedy)0.450.400.1800
B → D (optimal)0.400.850.3400

12. Beam search — width-k look-ahead

Section

Part 2 of 4

13. Beam search: maintain the top-k partial sequences

Concept

Beam search keeps k hypotheses ("beams") alive at every step. At each step, each beam is expanded to all V continuations, producing k·V candidates. The top-k by cumulative log-prob survive into the next step.

\[ \text{score}(w_{1:t}) = \sum_{i=1}^{t} \log P(w_i \mid w_{1:i-1}) \]

Log-probabilities are used (not raw probs) to avoid numerical underflow when multiplying many small numbers. k=1 recovers greedy; k=V is exact search (exponential cost). Typical NMT: k=4–8.

14. By analogy: Beam search: maintain the top-k partial sequences

Analogy

Discussion prompt

Explain Beam search: maintain the top-k partial sequences by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

Log-probabilities are used (not raw probs) to avoid numerical underflow when multiplying many small numbers. k=1 recovers greedy; k=V is exact search (exponential cost). Typical NMT: k=4–8.

15. What has to be given first: Beam search trace — k=2, 2 steps

Missing information

Discussion prompt

Vocab = [I, am, not, here]. Step 0: one initial hypothesis (empty). We expand to all 4 tokens, keep top-2. Step 1: expand each surviving beam to 4 tokens (8 candidates), keep top-2.

What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.

Hint: Anything you would have to invent to get started is a thing the problem must supply.

Answer:

After 'I', the model assigns high probability to 'am' (log-P step1 = -0.1050), giving total -0.6755 + (-0.1050) = -0.7805. This beats any continuation from 'am here' except its own top path.

16. Beam search trace — k=2, 2 steps

Worked example

Vocab = [I, am, not, here]. Step 0: one initial hypothesis (empty). We expand to all 4 tokens, keep top-2. Step 1: expand each surviving beam to 4 tokens (8 candidates), keep top-2.

import torch, torch.nn.functional as F
VOCAB = ["I","am","not","here"]
# step 0: one hypothesis -> expand to |VOCAB|
lp0 = F.log_softmax(torch.tensor([2.0,1.5,0.5,0.0]),dim=-1)
top2 = lp0.topk(2)
b_toks = top2.indices.tolist()   # [0,1] = I, am
b_lps  = top2.values.tolist()    # [-0.6755, -1.1755]
# step 1: expand each beam
lp1_I  = F.log_softmax(torch.tensor([-1.,3.5,0.5,0.5]),dim=-1)
lp1_am = F.log_softmax(torch.tensor([-1.,-1.,2.0,3.0]),dim=-1)
cands = []
for bi,(tok,blp) in enumerate(zip(b_toks,b_lps)):
    lp_nxt = lp1_I if bi==0 else lp1_am
    for vi in range(4):
        cands.append((VOCAB[tok],VOCAB[vi],blp+lp_nxt[vi].item()))
cands.sort(key=lambda x: x[2], reverse=True)
print(cands[:2])
stepbeam A (log-P)beam B (log-P)action
0 initI (-0.6755)am (-1.1755)top-2 from V=4
1 expandI am (-0.7805)am here (-1.5152)8 cands → top-2
winnerI amam herebest 2-token seqs

Beam A = 'I am', log-P = -0.7805; Beam B = 'am here', log-P = -1.5152

Why: After 'I', the model assigns high probability to 'am' (log-P step1 = -0.1050), giving total -0.6755 + (-0.1050) = -0.7805. This beats any continuation from 'am here' except its own top path.

17. Work backwards from the answer: Beam search trace — k=2, 2 steps

Reverse engineer

Discussion prompt

Work backwards. The example finished here:

Beam A = 'I am', log-P = -0.7805; Beam B = 'am here', log-P = -1.5152

What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.

Hint: Every quantity in the result had to enter somewhere. Account for each one.

Answer:

Vocab = [I, am, not, here]. Step 0: one initial hypothesis (empty). We expand to all 4 tokens, keep top-2. Step 1: expand each surviving beam to 4 tokens (8 candidates), keep top-2.

18. Length normalization

Concept

Raw beam scores are sums of log-probs. Adding more tokens always decreases the score (each log-prob is negative). Without correction, beam search systematically prefers shorter sequences.

\[ \text{score}_{\text{norm}}(w_{1:T}) = \frac{1}{\text{lp}(T)} \sum_{t=1}^{T} \log P(w_t \mid w_{<t}), \quad \text{lp}(T) = \frac{(5+T)^\alpha}{6^\alpha} \]

length Tlp(T) alpha=0.6raw log-Pnormalized score
21.097-1.200-1.094
51.451-2.800-1.930

Google NMT uses alpha=0.6. Alpha=0 gives no normalization; alpha=1 averages log-probs. In this example the length-5 sequence still loses after normalization — but the penalty is reduced so shorter isn't artificially favored.

19. Break it if you can: Length normalization

Counterexample

Discussion prompt

Raw beam scores are sums of log-probs. Adding more tokens always decreases the score (each log-prob is negative). Without correction, beam search systematically prefers shorter sequences.

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.

20. Something is wrong here: comparing beams by raw cumulative log-prob

Anomaly

Predict first

A student writes this, and it looks reasonable:

After decoding, rank final hypotheses by their total log-probability score.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: Raw log-prob penalizes every additional token.

Divide each beam's cumulative log-prob by its length penalty before comparing.

Why: Raw log-prob penalizes every additional token. A 2-token sequence will almost always outscore a 5-token sequence of equal per-token quality.

21. Trap: comparing beams by raw cumulative log-prob

Trap

The trap

After decoding, rank final hypotheses by their total log-probability score.

Sort beams: score(seq_A_len2) = -1.2 > score(seq_B_len5) = -2.8 → seq_A wins

Why: Raw log-prob penalizes every additional token. A 2-token sequence will almost always outscore a 5-token sequence of equal per-token quality.

The fix

Divide each beam's cumulative log-prob by its length penalty before comparing.

Normalize: seq_A = -1.2/1.097 = -1.094, seq_B = -2.8/1.451 = -1.930 → seq_A still wins, but fairly

Why: Now both sequences are compared on a per-token basis. Longer sequences are no longer systematically penalized; you can tune alpha to preference.

22. Sampling — temperature, top-k, top-p

Section

Part 3 of 4

23. Temperature scaling

Concept

Before applying softmax, divide every logit by temperature T. This rescales the distribution without changing the argmax order, but sharply changes the shape of the probability mass.

\[ P_T(w_i) = \frac{\exp(z_i / T)}{\sum_j \exp(z_j / T)} \]

TP(cat)P(sat)P(fat)entropy (nats)
0.50.70660.21280.00180.8519
1.00.46800.25690.02331.3963
2.00.31160.23080.06951.6728

T < 1 sharpens the distribution (lower entropy, more deterministic). T > 1 flattens it (higher entropy, more random). T → 0 recovers greedy; T → infinity recovers the uniform distribution.

24. Watch it run: Temperature scaling

Pattern

Step through it

Step through Temperature scaling one row at a time. What is driving the change, and what would the row after the last one be?

  1. Step 1: T is 0.5
  2. Step 2: T is 1.0
  3. Step 3: T is 2.0

25. Predict the next row: Implementing temperature sampling

Pattern

Predict first

The table runs: 0.5 | cat | 0.7066 | very likely — near-greedy · 1.0 | cat | 0.4680 | baseline distribution

In Implementing temperature sampling, given the rows so far: what is the next one — the row where T is 2.0?

Correct: 2.0 | sat | 0.2308 | less dominant — more variety

Tsampled tokenP(token)interpretation
0.5cat0.7066very likely — near-greedy
1.0cat0.4680baseline distribution
2.0sat0.2308less dominant — more variety

Why: The relationship between the columns, not the individual numbers, is what generates the next row. With seed 42, lower T concentrates mass on 'cat'; at T=2.0 the distribution is flat enough that 'sat' appears more often.

26. Implementing temperature sampling

Worked example

One-liner: divide logits by T before softmax. Then torch.multinomial samples from the resulting distribution. Note: T=0 is a singularity — use argmax directly for greedy.

import torch, torch.nn.functional as F
torch.manual_seed(42)
vocab = ["cat","sat","mat","bat","hat","fat"]
logits = torch.tensor([2.1, 1.5, 0.8, 0.3,-0.2,-0.9])
def temp_sample(logits, T=1.0):
    probs = F.softmax(logits / T, dim=-1)
    idx   = torch.multinomial(probs, num_samples=1).item()
    return vocab[idx], probs[idx].item()
for T in [0.5, 1.0, 2.0]:
    tok, p = temp_sample(logits, T)
    print(f"T={T}: sampled '{tok}' (p={p:.4f})")
Tsampled tokenP(token)interpretation
0.5cat0.7066very likely — near-greedy
1.0cat0.4680baseline distribution
2.0sat0.2308less dominant — more variety

Seed 42: T=0.5 samples 'cat' (p=0.7066), T=1.0 samples 'cat' (p=0.4680), T=2.0 samples 'sat' (p=0.2308)

Why: With seed 42, lower T concentrates mass on 'cat'; at T=2.0 the distribution is flat enough that 'sat' appears more often. These exact values come from torch.multinomial with manual_seed(42).

27. Top-k sampling

Concept

Top-k sampling restricts sampling to the k highest-probability tokens. All other logits are set to -inf before softmax, giving them probability 0. The remaining k tokens are renormalized and one is sampled.

\[ \tilde{z}_i = \begin{cases} z_i & \text{if } i \in \text{top-}k \\ -\infty & \text{otherwise} \end{cases}, \quad P_{\text{top-}k}(w_i) = \text{softmax}(\tilde{z})_i \]

tokenraw PP after top-3 maskin nucleus?
cat0.46800.5490yes
sat0.25690.3013yes
mat0.12750.1496yes
bat0.07740.0000no
hat0.04690.0000no
fat0.02330.0000no

k=3: bat, hat, fat are zeroed. The remaining three are renormalized so they sum to 1. Problem: a fixed k may include too many low-quality tokens when the distribution is flat, or too few when it is sharp.

28. Which is which, by in nucleus?

Discrimination

Sort into buckets

Sort these by in nucleus?, from memory, without looking back at Top-k sampling. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

yes
cat; sat; mat
no
bat; hat; fat
g1
in nucleus? is "yes" for cat, sat, mat — that is what the table on "Top-k sampling" records, and it is the single property separating this group from the rest.
g2
in nucleus? is "no" for bat, hat, fat — that is what the table on "Top-k sampling" records, and it is the single property separating this group from the rest.

29. Top-p (nucleus) sampling

Concept

Top-p (nucleus) sampling uses a dynamic cutoff. Sort tokens by descending probability; take the smallest prefix whose cumulative probability meets or exceeds p. Sample uniformly from this nucleus.

\[ \mathcal{N}_p = \min_{S \subseteq V} S \quad\text{s.t.}\quad \sum_{w \in S} P(w) \geq p \]

tokenP (sorted desc)cumulative Pin nucleus p=0.90?
cat0.46800.4680yes
sat0.25690.7249yes
mat0.12750.8524yes
bat0.07740.9298 >=0.90yes (cutoff here)
hat0.04690.9767no
fat0.02331.0000no

p=0.90: nucleus = {cat, sat, mat, bat}, renormalized probs [0.5034, 0.2763, 0.1372, 0.0832]. When the model is confident (sharp distribution) the nucleus shrinks; when the model is uncertain it grows — adapting automatically.

30. Fill in: cumulative P for Top-p (nucleus) sampling

Comparison

Comparison matrix

From Top-p (nucleus) sampling: refill the cumulative P column from what you know. The rest of the table is as it appeared.

tokenP (sorted desc)cumulative Pin nucleus p=0.90?
cat0.46800.4680yes
sat0.25690.7249yes
mat0.12750.8524yes
bat0.07740.9298 >=0.90yes (cutoff here)
hat0.04690.9767no
fat0.02331.0000no

31. Guess the shape of the answer: Implementing top-p from scratch

Estimation

Predict first

Standard recipe: sort descending, compute cumulative sum, find the first index where cumsum >= p, zero everything after, renormalize, and sample. All operations are differentiable-friendly (no Python loops over V needed).

Commit before you compute: what does Implementing top-p from scratch come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: Nucleus = {cat, sat, mat, bat}; renorm probs = [0.5034, 0.2763, 0.1372, 0.0832]; sample → 'cat'

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The remove mask uses cumsum - sorted_p > p: for 'hat', cumsum=0.9767 but cumsum - P(hat) = 0.9767 - 0.0469 = 0.9298 > 0.90, so hat is excluded.

32. Implementing top-p from scratch

Worked example

Standard recipe: sort descending, compute cumulative sum, find the first index where cumsum >= p, zero everything after, renormalize, and sample. All operations are differentiable-friendly (no Python loops over V needed).

import torch, torch.nn.functional as F
torch.manual_seed(0)
vocab = ["cat","sat","mat","bat","hat","fat"]
logits = torch.tensor([2.1,1.5,0.8,0.3,-0.2,-0.9])
def nucleus_sample(logits, p=0.90):
    probs = F.softmax(logits, dim=-1)
    sorted_p, sorted_idx = probs.sort(descending=True)
    cumsum = sorted_p.cumsum(dim=-1)
    # keep up to (and including) first index where cumsum >= p
    remove = cumsum - sorted_p > p
    sorted_p[remove] = 0.0
    sorted_p /= sorted_p.sum()  # renormalize
    sample_rank = torch.multinomial(sorted_p, 1).item()
    return vocab[sorted_idx[sample_rank].item()]
print(nucleus_sample(logits, p=0.90))  # 'cat'
token (sorted)raw Pcumsumremove? (cumsum-P > 0.90)
cat0.46800.4680no
sat0.25690.7249no
mat0.12750.8524no
bat0.07740.9298no
hat0.04690.9767yes
fat0.02331.0000yes

Nucleus = {cat, sat, mat, bat}; renorm probs = [0.5034, 0.2763, 0.1372, 0.0832]; sample → 'cat'

Why: The remove mask uses cumsum - sorted_p > p: for 'hat', cumsum=0.9767 but cumsum - P(hat) = 0.9767 - 0.0469 = 0.9298 > 0.90, so hat is excluded. This one-liner correctly keeps the cutoff token itself.

33. Which is which, by remove? (cumsum-P > 0.90)

Discrimination

Sort into buckets

Sort these by remove? (cumsum-P > 0.90), from memory, without looking back at Implementing top-p from scratch. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

no
cat; sat; mat; bat
yes
hat; fat
g1
remove? (cumsum-P > 0.90) is "no" for cat, sat, mat, bat — that is what the table on "Implementing top-p from scratch" records, and it is the single property separating this group from the rest.
g2
remove? (cumsum-P > 0.90) is "yes" for hat, fat — that is what the table on "Implementing top-p from scratch" records, and it is the single property separating this group from the rest.

34. Something is wrong here: confusing top-k nucleus size with top-p nucleus size

Anomaly

Predict first

A student writes this, and it looks reasonable:

Top-k=4 and top-p=0.90 always produce the same 4-token nucleus for these logits.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: Coincidentally, p=0.90 includes exactly 4 tokens here.

The nucleus size for top-p is data-dependent; top-k is always exactly k.

Why: Coincidentally, p=0.90 includes exactly 4 tokens here. But this is a numerical coincidence, not a rule.

35. Trap: confusing top-k nucleus size with top-p nucleus size

Trap

The trap

Top-k=4 and top-p=0.90 always produce the same 4-token nucleus for these logits.

Both include {cat, sat, mat, bat} — so they are equivalent

Why: Coincidentally, p=0.90 includes exactly 4 tokens here. But this is a numerical coincidence, not a rule.

The fix

The nucleus size for top-p is data-dependent; top-k is always exactly k.

If the model is very confident (P(cat)=0.95), nucleus p=0.90 shrinks to 1 token; top-k=4 still includes 4 tokens regardless

Why: This is the key motivation for top-p: when the model is confident, restrict sampling tightly; when uncertain, allow more diversity. Top-k cannot adapt this way.

36. Break it on purpose: confusing top-k nucleus size with top-p…

Break the constraint

Discussion prompt

The rule this trap just fixed:

The nucleus size for top-p is data-dependent; top-k is always exactly k.

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:

Coincidentally, p=0.90 includes exactly 4 tokens here. But this is a numerical coincidence, not a rule.

37. Decoder selection and the full pipeline

Section

Part 4 of 4

38. Restore the missing line: Beam search from scratch — full Python

Fill the middle

Fill in the blanks

From Beam search from scratch — full Python — one line has had its right-hand side removed. Put it back.

import torch, torch.nn.functional as F
def beam_search(logit_fn, bos_id, eos_id, k=2, max_len=5):
"""logit_fn(seq) -> logit tensor of shape [V]"""
beams = [([bos_id], 0.0)] # (tokens, log-prob)
for _ in range(max_len):
cands = []
for toks, lp in beams:
if toks[-1] == eos_id:
cands.append((toks, lp)); continue
logits = logit_fn(toks)
lp_next = F.log_softmax(logits, dim=-1)
for v in range(len(logits)):
cands.append((toks + [v], lp + lp_next[v].item()))
cands.sort(key=lambda x: x[1], reverse=True)
beams = cands[:k]
return beams
print("beam_search defined — attach a real logit_fn to run")

Why: lp_next is what everything below it consumes, so the wrong expression here fails later and somewhere else. Total cost is O(k · V · T) where T is the max decode length.

39. Beam search from scratch — full Python

Concept

A clean beam-search loop: maintain a list of (sequence, cumulative_log_prob) pairs; at each step expand every hypothesis, flatten to k·V candidates, keep top-k. Terminate when all beams emit <eos> or reach max length.

import torch, torch.nn.functional as F
def beam_search(logit_fn, bos_id, eos_id, k=2, max_len=5):
    """logit_fn(seq) -> logit tensor of shape [V]"""
    beams = [([bos_id], 0.0)]  # (tokens, log-prob)
    for _ in range(max_len):
        cands = []
        for toks, lp in beams:
            if toks[-1] == eos_id:
                cands.append((toks, lp)); continue
            logits = logit_fn(toks)
            lp_next = F.log_softmax(logits, dim=-1)
            for v in range(len(logits)):
                cands.append((toks + [v], lp + lp_next[v].item()))
        cands.sort(key=lambda x: x[1], reverse=True)
        beams = cands[:k]
    return beams
print("beam_search defined — attach a real logit_fn to run")
iterationbeams alivecandidates generatedkept
0 (init)1Vk=2
122Vk=2
tkk*Vk=2

At each step: O(k · V) log-softmax evaluations → sort → keep top-k

Why: Total cost is O(k · V · T) where T is the max decode length. k=1 reduces to greedy; k=V is exact but has exponential memory.

40. What stays fixed: Beam search from scratch — full Python

Invariant

Step through it

Step through Beam search from scratch — full Python one row at a time. One of these columns never changes — find it, and say why it cannot.

  1. Step 1: iteration is 0 (init)
  2. Step 2: iteration is 1
  3. Step 3: iteration is t

41. Without one step: Decoding strategy selection — the recipe

Constraint

Discussion prompt

Run Decoding strategy selection — the recipe with this step confiscated:

Add a truncation filter. Prefer top-p (p=0.90–0.95) over top-k because it adapts to the model's confidence. Use top-k as a hard cap if needed (k=40–100 is common).

Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.

Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.

Answer:

  1. Decide on determinism vs diversity. Translation, summarization, code → near-deterministic → use beam search (k=4–8) with length normalization. Creative…
  2. Pick a base sampler. Start with temperature T=1.0 (baseline). Lower T (0.7–0.9) for more focused outputs; raise T (1.1–1.5) for more variety.
  3. Add a truncation filter. Prefer top-p (p=0.90–0.95) over top-k because it adapts to the model's confidence. Use top-k as a hard cap if needed (k=40–100 is…
  4. Apply length penalty if using beam search. Google NMT alpha=0.6 is a safe default; tune on your validation set.
  5. Combine carefully. Temperature + top-p is standard (apply T first, then compute nucleus on scaled probs). Avoid beam search + temperature: beam uses…

42. Decoding strategy selection — the recipe

Pattern

  1. Decide on determinism vs diversity. Translation, summarization, code → near-deterministic → use beam search (k=4–8) with length normalization. Creative writing, chat → stochastic → use sampling.
  2. Pick a base sampler. Start with temperature T=1.0 (baseline). Lower T (0.7–0.9) for more focused outputs; raise T (1.1–1.5) for more variety.
  3. Add a truncation filter. Prefer top-p (p=0.90–0.95) over top-k because it adapts to the model's confidence. Use top-k as a hard cap if needed (k=40–100 is common).
  4. Apply length penalty if using beam search. Google NMT alpha=0.6 is a safe default; tune on your validation set.
  5. Combine carefully. Temperature + top-p is standard (apply T first, then compute nucleus on scaled probs). Avoid beam search + temperature: beam uses deterministic ranking, stochasticity undermines it.

43. Where does it stop working: Decoding strategy selection — the recipe

Edge cases

Discussion prompt

Decoding strategy selection — the 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. Decide on determinism vs diversity. Translation, summarization, code → near-deterministic → use beam search (k=4–8) with length normalization. Creative…
  2. Pick a base sampler. Start with temperature T=1.0 (baseline). Lower T (0.7–0.9) for more focused outputs; raise T (1.1–1.5) for more variety.
  3. Add a truncation filter. Prefer top-p (p=0.90–0.95) over top-k because it adapts to the model's confidence. Use top-k as a hard cap if needed (k=40–100 is…
  4. Apply length penalty if using beam search. Google NMT alpha=0.6 is a safe default; tune on your validation set.
  5. Combine carefully. Temperature + top-p is standard (apply T first, then compute nucleus on scaled probs). Avoid beam search + temperature: beam uses…

44. Rule out three: Check 1 — Greedy suboptimality

Elimination

Eliminate the wrong options

Vocab = [A, B]. Step-0 probs: P(A)=0.45, P(B)=0.40. Step-1 probs after A: P(D|A)=0.40 (best continuation). Step-1 probs after B: P(D|B)=0.85 (best continuation). Greedy picks A at step 0. What is greedy's best 2-token joint probability, and does it equal the globally optimal sequence's joint probability?

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. Greedy best = 0.45×0.40 = 0.180; optimal is also A→D with 0.180
  • B. Greedy best = 0.45×0.40 = 0.180; optimal is B→D with 0.340 — greedy is suboptimal
  • C. Greedy best = 0.45×0.85 = 0.383; optimal is A→D with 0.383
  • D. Greedy cannot be compared to optimal on a 2-token sequence — suboptimality only manifests at length ≥ 5

Survives elimination: B

Why: Greedy picks A (P=0.45 > 0.40). The best continuation from A is D with P(D|A)=0.40, giving joint 0.45×0.40=0.180. Meanwhile, B→D has joint 0.40×0.85=0.340. So greedy's choice locks in a suboptimal path: 0.340 > 0.180. This is the exact counter-example from slide 6 of this deck.

45. Check 1 — Greedy suboptimality

Check

Work this through before clicking.

Check your understanding

Vocab = [A, B]. Step-0 probs: P(A)=0.45, P(B)=0.40. Step-1 probs after A: P(D|A)=0.40 (best continuation). Step-1 probs after B: P(D|B)=0.85 (best continuation). Greedy picks A at step 0. What is greedy's best 2-token joint probability, and does it equal the globally optimal sequence's joint probability?

  • A. Greedy best = 0.45×0.40 = 0.180; optimal is also A→D with 0.180
  • B. Greedy best = 0.45×0.40 = 0.180; optimal is B→D with 0.340 — greedy is suboptimal (correct)
  • C. Greedy best = 0.45×0.85 = 0.383; optimal is A→D with 0.383
  • D. Greedy cannot be compared to optimal on a 2-token sequence — suboptimality only manifests at length ≥ 5

Answer: B

Why: Greedy picks A (P=0.45 > 0.40). The best continuation from A is D with P(D|A)=0.40, giving joint 0.45×0.40=0.180. Meanwhile, B→D has joint 0.40×0.85=0.340. So greedy's choice locks in a suboptimal path: 0.340 > 0.180. This is the exact counter-example from slide 6 of this deck.

Why A tempts people
This misses that B→D achieves 0.340. Greedy's path A→D is 0.180, which is not optimal — B→D nearly doubles the joint probability.
Why C tempts people
P(D|A)=0.40, not 0.85. P(D|B)=0.85 is the continuation after B, not after A. Mixing up the conditional probabilities is a classic distractor on beam-search questions.
Why D tempts people
Greedy suboptimality can manifest in sequences as short as 2 tokens, as this example proves. There is no minimum length threshold — the cliff phenomenon can occur at any step.

46. Answer it before you see the options: Check 2 — Temperature and entropy

Prediction

Predict first

For logits [2.1, 1.5, 0.8, 0.3, -0.2, -0.9], the entropies at T=0.5, 1.0, 2.0 are 0.852, 1.396, 1.673 nats respectively. Which claim follows directly?

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: T=2.0 produces higher entropy than T=1.0, so it samples more uniformly across tokens

Why: H(T=2.0)=1.673 > H(T=1.0)=1.396 > H(T=0.5)=0.852. Higher entropy means the distribution is more spread out, so T=2.0 spreads probability more uniformly — statement A is correct.

47. Check 2 — Temperature and entropy

Check

Entropy H = -sum(p * log p). Use the verified numbers from the temperature table.

Check your understanding

For logits [2.1, 1.5, 0.8, 0.3, -0.2, -0.9], the entropies at T=0.5, 1.0, 2.0 are 0.852, 1.396, 1.673 nats respectively. Which claim follows directly?

  • A. T=2.0 produces higher entropy than T=1.0, so it samples more uniformly across tokens (correct)
  • B. T=0.5 is equivalent to beam search with k=1
  • C. Temperature does not change which token has the highest probability
  • D. At T→0, every token has nonzero probability

Answer: A

Why: H(T=2.0)=1.673 > H(T=1.0)=1.396 > H(T=0.5)=0.852. Higher entropy means the distribution is more spread out, so T=2.0 spreads probability more uniformly — statement A is correct.

Why B tempts people
Temperature sampling is still stochastic (multinomial draw) even at T=0.5. Beam search k=1 (greedy) is deterministic argmax. They are not equivalent — with seed 42 T=0.5 sampled 'cat' with p=0.7066, which happens to match greedy, but a different seed could sample a non-top token.
Why C tempts people
This is actually true in general (dividing all logits by T preserves their ordering), but it is NOT what 'follows directly' from the entropy data alone, and it is not the most precise or exam-relevant reading of the table. The question asks what 'follows directly' from the entropy numbers.
Why D tempts people
At T→0, logits/T→±∞. The argmax token's probability approaches 1 and all others approach 0 (not nonzero). This is the opposite of D.

48. Rule out three: Check 3 — Top-p vs top-k nucleus size

Elimination

Eliminate the wrong options

For the distribution above, the top-p nucleus with p=0.85 contains how many tokens?

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. 2 tokens (cumsum reaches 0.725 at index 1)
  • B. 3 tokens (cumsum reaches 0.852 at index 2, first value ≥ 0.85)
  • C. 4 tokens (same as top-k with k=4)
  • D. 3 tokens but only if p=0.90, not p=0.85

Survives elimination: B

Why: The nucleus is the smallest prefix whose cumulative probability >= p=0.85. cumsum[0]=0.468 < 0.85; cumsum[1]=0.725 < 0.85; cumsum[2]=0.852 >= 0.85. So the nucleus includes the first 3 tokens (indices 0, 1, 2 in sorted order: cat, sat, mat).

49. Check 3 — Top-p vs top-k nucleus size

Check

Use the verified cumsum table: sorted probs [0.4680, 0.2569, 0.1275, 0.0774, 0.0469, 0.0233], cumsum [0.468, 0.725, 0.852, 0.930, 0.977, 1.000].

Check your understanding

For the distribution above, the top-p nucleus with p=0.85 contains how many tokens?

  • A. 2 tokens (cumsum reaches 0.725 at index 1)
  • B. 3 tokens (cumsum reaches 0.852 at index 2, first value ≥ 0.85) (correct)
  • C. 4 tokens (same as top-k with k=4)
  • D. 3 tokens but only if p=0.90, not p=0.85

Answer: B

Why: The nucleus is the smallest prefix whose cumulative probability >= p=0.85. cumsum[0]=0.468 < 0.85; cumsum[1]=0.725 < 0.85; cumsum[2]=0.852 >= 0.85. So the nucleus includes the first 3 tokens (indices 0, 1, 2 in sorted order: cat, sat, mat).

Why A tempts people
cumsum[1]=0.725 does not yet reach 0.85, so the nucleus is not complete after 2 tokens. We must keep adding tokens until cumsum >= p.
Why C tempts people
4 tokens is the nucleus for p=0.90 (cumsum[3]=0.930 is the first value ≥ 0.90). For p=0.85, the cutoff occurs earlier at index 2 (cumsum=0.852).
Why D tempts people
The 3-token nucleus applies at p=0.85 (cutoff at cumsum=0.852). At p=0.90 the nucleus grows to 4 tokens (cutoff at cumsum=0.930). D has the threshold backwards.

50. Task-decoder pairing — when to use what

Concept

taskrecommended decodertypical settings
Machine translationBeam search + length normk=4–8, alpha=0.6
SummarizationBeam searchk=4, repetition penalty
Code completionBeam or greedyT=0.2 or greedy
Open-ended chatTop-p + temperatureT=0.9, p=0.92
Creative writingTop-p + high TT=1.1–1.5, p=0.95
Constrained output (JSON)Constrained beamlogit masking

The intuition: high-stakes or single-correct-answer tasks want near-determinism (beam); open-ended tasks want diversity (sampling). Code is unusual — it benefits from creativity in naming but correctness in logic, so low T or constrained beam is preferred.

51. Your turn — project brief

Concept

You will build a decoding sandbox: a self-contained Python module that wraps a dummy bigram language model and exposes greedy, beam search (k), temperature+top-p, and length-normalized scoring — then compares them head-to-head on a fixed sequence.

Milestone 1
Build a tiny bigram LM with a fixed transition matrix. Implement greedy decode.
Milestone 2
Add beam search (k configurable). Verify trace table matches manual log-prob arithmetic.
Milestone 3
Add temperature + top-p sampling. Run 100 samples at each T, plot token frequency.
Milestone 4
Add length normalization. Compare raw vs normalized beam ranking on sequences of different lengths.

52. Which is which: Your turn — project brief

Matching

Match the pairs

From Your turn — project brief — match each one to what it actually does. The descriptions have been shuffled.

  • c1. Milestone 1
  • c2. Milestone 2
  • c3. Milestone 3
  • c4. Milestone 4
  • b1. Build a tiny bigram LM with a fixed transition matrix. Implement greedy decode.
  • b2. Add beam search (k configurable). Verify trace table matches manual log-prob arithmetic.
  • b3. Add temperature + top-p sampling. Run 100 samples at each T, plot token frequency.
  • b4. Add length normalization. Compare raw vs normalized beam ranking on sequences of different lengths.

Why: Milestone 1, Milestone 2, Milestone 3, Milestone 4 are easy to tell apart while they are sitting next to their descriptions and much harder afterwards, which is what this checks.

53. Predict the next row: Milestone 1 — Greedy on a bigram LM

Pattern

Predict first

The table runs: 0 | <s> | computed by argmax(W[0]) | max of softmax(W[0]) · 1 | greedy_0 | computed by argmax(W[greedy_0]) | max of softmax(W[greedy_0])

In Milestone 1 — Greedy on a bigram LM, given the rows so far: what is the next one — the row where step is ...?

Correct: ... | ... | ... | ...

stepcurrent tokengreedy next tokenP(next)
0<s>computed by argmax(W[0])max of softmax(W[0])
1greedy_0computed by argmax(W[greedy_0])max of softmax(W[greedy_0])
............

Why: The relationship between the columns, not the individual numbers, is what generates the next row. The transition matrix W is the entire model.

54. Milestone 1 — Greedy on a bigram LM

Worked example

Build a tiny bigram LM as a fixed [V, V] transition matrix (row = current token, column = next token logits). Greedy decode from a start token.

import torch, torch.nn.functional as F
torch.manual_seed(0)
VOCAB = ["<s>","I","am","here","<e>"]
V = len(VOCAB)
# random transition logits (fixed seed -> deterministic)
W = torch.randn(V, V)
def greedy_decode(start=0, max_len=6):
    seq = [start]; tok = start
    for _ in range(max_len):
        nxt = F.softmax(W[tok], dim=-1).argmax().item()
        seq.append(nxt)
        if nxt == V-1: break
        tok = nxt
    return [VOCAB[i] for i in seq]
print(greedy_decode())
stepcurrent tokengreedy next tokenP(next)
0<s>computed by argmax(W[0])max of softmax(W[0])
1greedy_0computed by argmax(W[greedy_0])max of softmax(W[greedy_0])
............

Run and print the decoded token list; verify each step by checking that F.softmax(W[tok]).argmax() == next token

Why: The transition matrix W is the entire model. torch.manual_seed(0) makes W deterministic so your output is reproducible. This tiny LM has no meaning — it is a structural scaffold for testing decoders.

55. Work backwards from the answer: Milestone 1 — Greedy on a bigram LM

Reverse engineer

Discussion prompt

Work backwards. The example finished here:

Run and print the decoded token list; verify each step by checking that F.softmax(W[tok]).argmax() == next token

What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.

Hint: Every quantity in the result had to enter somewhere. Account for each one.

Answer:

Build a tiny bigram LM as a fixed [V, V] transition matrix (row = current token, column = next token logits). Greedy decode from a start token.

56. What has to be given first: Milestone 2 — Beam search on the bigram LM

Missing information

Discussion prompt

Plug the same bigram LM into the beam_search function from slide 17. Set logit_fn = lambda seq: W[seq[-1]]. Print the top-k beams and their log-probs after 3 steps.

What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.

Hint: Anything you would have to invent to get started is a thing the problem must supply.

Answer:

If your beam_decode log-prob matches the manual sum, the implementation is correct. This is the trace-table check the exam will expect you to perform.

57. Milestone 2 — Beam search on the bigram LM

Worked example

Plug the same bigram LM into the beam_search function from slide 17. Set logit_fn = lambda seq: W[seq[-1]]. Print the top-k beams and their log-probs after 3 steps.

# Continuing from milestone 1 (W and VOCAB defined)
import torch, torch.nn.functional as F
torch.manual_seed(0)
VOCAB = ["<s>","I","am","here","<e>"]
V = len(VOCAB)
W = torch.randn(V, V)
def beam_decode(start=0, k=2, max_len=3):
    beams = [([start], 0.0)]
    for _ in range(max_len):
        cands = []
        for toks, lp in beams:
            lp_nxt = F.log_softmax(W[toks[-1]], dim=-1)
            for v in range(V):
                cands.append((toks+[v], lp+lp_nxt[v].item()))
        cands.sort(key=lambda x: x[1], reverse=True)
        beams = cands[:k]
    return [([ VOCAB[t] for t in toks], round(lp,4)) for toks,lp in beams]
print(beam_decode(k=2))
beam ranksequencecumulative log-prob
1top-1 beam (printed)highest score
2top-2 beam (printed)second highest score

Verify: manually compute log-prob of the top beam by summing F.log_softmax(W[tok])[next_tok] at each step

Why: If your beam_decode log-prob matches the manual sum, the implementation is correct. This is the trace-table check the exam will expect you to perform.

58. What each one costs: Milestone 2 — Beam search on the bigram LM

Trade off

Comparison matrix

From Milestone 2 — Beam search on the bigram LM: every row here is a choice with a cost. Fill the cumulative log-prob column, then say which row you would actually pick and what you give up for it.

beam ranksequencecumulative log-prob
1top-1 beam (printed)highest score
2top-2 beam (printed)second highest score

59. Guess the shape of the answer: Show it off — full comparison

Estimation

Predict first

Run all four decoders on the same bigram LM and print a side-by-side summary table. Include: sequence, joint log-prob, length-normalized score (alpha=0.6).

Commit before you compute: what does Show it off — full comparison come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: Expected: beam k=2 normalized score >= greedy normalized score in most cases, especially when the bigram LM has high-probability long paths

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Length normalization prevents beam search from always preferring shorter sequences.

60. Show it off — full comparison

Worked example

Run all four decoders on the same bigram LM and print a side-by-side summary table. Include: sequence, joint log-prob, length-normalized score (alpha=0.6).

import torch, torch.nn.functional as F
torch.manual_seed(0)
VOCAB = ["<s>","I","am","here","<e>"]; V = len(VOCAB)
W = torch.randn(V, V)
def lp_norm(lp, T, alpha=0.6):
    return lp / (((5+T)/6)**alpha)
# greedy
seq=[0]; tok=0
for _ in range(4):
    nxt=F.softmax(W[tok],dim=-1).argmax().item()
    seq.append(nxt)
    if nxt==V-1: break
    tok=nxt
g_lp=sum(F.log_softmax(W[seq[i]],dim=-1)[seq[i+1]].item() for i in range(len(seq)-1))
print("Greedy:", [VOCAB[t] for t in seq], f"lp={g_lp:.3f}", f"norm={lp_norm(g_lp,len(seq)-1):.3f}")
# beam k=2
def beam(k=2,ml=4):
    bms=[([0],0.)]
    for _ in range(ml):
        c=[(t+[v],lp+F.log_softmax(W[t[-1]],dim=-1)[v].item()) for t,lp in bms for v in range(V)]
        c.sort(key=lambda x:x[1],reverse=True); bms=c[:k]
    return bms
for s,lp in beam():
    print("Beam:", [VOCAB[t] for t in s], f"lp={lp:.3f}", f"norm={lp_norm(lp,len(s)-1):.3f}")
decodersequencelog-probnormalized score
greedy(your output)(printed)(printed)
beam k=2 top-1(your output)(printed)(printed)
beam k=2 top-2(your output)(printed)(printed)

Expected: beam k=2 normalized score >= greedy normalized score in most cases, especially when the bigram LM has high-probability long paths

Why: Length normalization prevents beam search from always preferring shorter sequences. If greedy outputs a shorter path, beam's longer top sequence may outscore it after normalization.

61. Fill in: log-prob for Show it off — full comparison

Comparison

Comparison matrix

From Show it off — full comparison: refill the log-prob column from what you know. The rest of the table is as it appeared.

decodersequencelog-probnormalized score
greedy(your output)(printed)(printed)
beam k=2 top-1(your output)(printed)(printed)
beam k=2 top-2(your output)(printed)(printed)

62. Connect it up: Lesson 104: Text Decoding Strategies

Connect it up

Draw it

One page, no notation unless you need it: draw how these connect — Greedy decoding — fast but fragile · Beam search — width-k look-ahead · Sampling — temperature, top-k, top-p · Decoder selection and the full pipeline. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.

63. Lesson 104 recap

Recap

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 104 — Greedy / Beam / Sampling Decoding — Barron · USAAIO Round 2 Preparation, 2026
  2. Wu et al. 'Google's Neural Machine Translation System: Bridging the Gap between Human and Machine Translation' (length penalty) — arXiv:1609.08144, 2016
  3. Holtzman et al. 'The Curious Case of Neural Text Degeneration' (top-p / nucleus sampling) — ICLR 2020, arXiv:1904.09751
  4. All log-softmax arithmetic, beam-search trace, temperature/top-k/top-p numbers verified with torch 2.7.1+cpu and numpy 2.2.6, June 2026 — Real execution, verified

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

Book on Wyzant · Text (657) 465-8108