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
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.
Objectives
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.
Section
Part 1 of 4
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.
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.
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.
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)| step | greedy token | P(token) | cumulative P |
|---|---|---|---|
| 0 | cat | 0.4680 | 0.4680 |
| 1 | sat | 0.7098 | 0.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.
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.
| step | greedy token | P(token) | cumulative P |
|---|---|---|---|
| 0 | cat | 0.4680 | 0.4680 |
| 1 | sat | 0.7098 | 0.3322 |
| — | sequence | — | 0.3322 |
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.
| path | P(tok 0) | P(tok 1 | tok 0) | joint P |
|---|---|---|---|
| A → D (greedy) | 0.45 | 0.40 | 0.1800 |
| B → D (optimal) | 0.40 | 0.85 | 0.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.
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.
| path | P(tok 0) | P(tok 1 | tok 0) | joint P |
|---|---|---|---|
| A → D (greedy) | 0.45 | 0.40 | 0.1800 |
| B → D (optimal) | 0.40 | 0.85 | 0.3400 |
Section
Part 2 of 4
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.
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.
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.
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])| step | beam A (log-P) | beam B (log-P) | action |
|---|---|---|---|
| 0 init | I (-0.6755) | am (-1.1755) | top-2 from V=4 |
| 1 expand | I am (-0.7805) | am here (-1.5152) | 8 cands → top-2 |
| winner | I am | am here | best 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.
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.
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 T | lp(T) alpha=0.6 | raw log-P | normalized score |
|---|---|---|---|
| 2 | 1.097 | -1.200 | -1.094 |
| 5 | 1.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.
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.
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.
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.
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.
Section
Part 3 of 4
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)} \]
| T | P(cat) | P(sat) | P(fat) | entropy (nats) |
|---|---|---|---|---|
| 0.5 | 0.7066 | 0.2128 | 0.0018 | 0.8519 |
| 1.0 | 0.4680 | 0.2569 | 0.0233 | 1.3963 |
| 2.0 | 0.3116 | 0.2308 | 0.0695 | 1.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.
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?
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
| T | sampled token | P(token) | interpretation |
|---|---|---|---|
| 0.5 | cat | 0.7066 | very likely — near-greedy |
| 1.0 | cat | 0.4680 | baseline distribution |
| 2.0 | sat | 0.2308 | less 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.
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})")| T | sampled token | P(token) | interpretation |
|---|---|---|---|
| 0.5 | cat | 0.7066 | very likely — near-greedy |
| 1.0 | cat | 0.4680 | baseline distribution |
| 2.0 | sat | 0.2308 | less 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).
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 \]
| token | raw P | P after top-3 mask | in nucleus? |
|---|---|---|---|
| cat | 0.4680 | 0.5490 | yes |
| sat | 0.2569 | 0.3013 | yes |
| mat | 0.1275 | 0.1496 | yes |
| bat | 0.0774 | 0.0000 | no |
| hat | 0.0469 | 0.0000 | no |
| fat | 0.0233 | 0.0000 | no |
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.
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.
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 \]
| token | P (sorted desc) | cumulative P | in nucleus p=0.90? |
|---|---|---|---|
| cat | 0.4680 | 0.4680 | yes |
| sat | 0.2569 | 0.7249 | yes |
| mat | 0.1275 | 0.8524 | yes |
| bat | 0.0774 | 0.9298 >=0.90 | yes (cutoff here) |
| hat | 0.0469 | 0.9767 | no |
| fat | 0.0233 | 1.0000 | no |
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.
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.
| token | P (sorted desc) | cumulative P | in nucleus p=0.90? |
|---|---|---|---|
| cat | 0.4680 | 0.4680 | yes |
| sat | 0.2569 | 0.7249 | yes |
| mat | 0.1275 | 0.8524 | yes |
| bat | 0.0774 | 0.9298 >=0.90 | yes (cutoff here) |
| hat | 0.0469 | 0.9767 | no |
| fat | 0.0233 | 1.0000 | no |
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.
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 P | cumsum | remove? (cumsum-P > 0.90) |
|---|---|---|---|
| cat | 0.4680 | 0.4680 | no |
| sat | 0.2569 | 0.7249 | no |
| mat | 0.1275 | 0.8524 | no |
| bat | 0.0774 | 0.9298 | no |
| hat | 0.0469 | 0.9767 | yes |
| fat | 0.0233 | 1.0000 | yes |
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.
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.
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.
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 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.
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.
Section
Part 4 of 4
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.
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")| iteration | beams alive | candidates generated | kept |
|---|---|---|---|
| 0 (init) | 1 | V | k=2 |
| 1 | 2 | 2V | k=2 |
| t | k | k*V | k=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.
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.
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:
Pattern
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:
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.
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.
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?
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.
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.
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?
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.
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.
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).
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?
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).
Concept
| task | recommended decoder | typical settings |
|---|---|---|
| Machine translation | Beam search + length norm | k=4–8, alpha=0.6 |
| Summarization | Beam search | k=4, repetition penalty |
| Code completion | Beam or greedy | T=0.2 or greedy |
| Open-ended chat | Top-p + temperature | T=0.9, p=0.92 |
| Creative writing | Top-p + high T | T=1.1–1.5, p=0.95 |
| Constrained output (JSON) | Constrained beam | logit 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.
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.
Matching
Match the pairs
From Your turn — project brief — match each one to what it actually does. The descriptions have been shuffled.
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.
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: ... | ... | ... | ...
| step | current token | greedy next token | P(next) |
|---|---|---|---|
| 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]) |
| ... | ... | ... | ... |
Why: The relationship between the columns, not the individual numbers, is what generates the next row. The transition matrix W is the entire model.
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())| step | current token | greedy next token | P(next) |
|---|---|---|---|
| 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]) |
| ... | ... | ... | ... |
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.
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.
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.
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 rank | sequence | cumulative log-prob |
|---|---|---|
| 1 | top-1 beam (printed) | highest score |
| 2 | top-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.
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 rank | sequence | cumulative log-prob |
|---|---|---|
| 1 | top-1 beam (printed) | highest score |
| 2 | top-2 beam (printed) | second highest score |
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.
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}")| decoder | sequence | log-prob | normalized 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.
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.
| decoder | sequence | log-prob | normalized score |
|---|---|---|---|
| greedy | (your output) | (printed) | (printed) |
| beam k=2 top-1 | (your output) | (printed) | (printed) |
| beam k=2 top-2 | (your output) | (printed) | (printed) |
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.
Recap
argmax at every step — O(V·T), deterministic, not globally optimal (proved by counter-example: A→D joint 0.18 < B→D joint 0.34)((5+T)/6)^alpha (alpha=0.6 default) — prevents beam search from preferring short sequencesWant this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.