Skip to content

Beam search, generic over a step function

ModuleL4.4 · build · Python · Pass 4 · 3 to 4 h, plus your graded tests (rung R5)
You buildpython/tinyllm/infer/beam.py: Hypothesis, select_state, beam_search, greedy_decode; and your own oracle tests in python/tests/l4-4-beam/
Contractcourse/contracts/py/tinyllm/infer/beam.pyi
Testscourse/tests/L4.4/test_beam.py (what they check: section 4); the oracle is exhaustive enumeration of every output of a toy model; your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsM09.2 log_softmax · L4.1 Seq2Seq.decode_step and L4.3 LuongAttention (the real decoder the tests search over) · reading: S-M06a (counting sequences) (or --ref-deps)
Used bylater: L6.7 the zoo’s beam decoding of seq2seq and transformer checkpoints (joins the registry with B7) · later: L5.5
MilestoneMS-L4 ({tinyllm} translate --beam 5 must beat greedy)
Optional depthGraves, “Sequence Transduction with Recurrent Neural Networks” (2012), section 3.2; Wu et al., “Google’s Neural Machine Translation System” (2016), section 7 (length penalty); Meister, Cotterell, and Vieira, “If beam search is the answer, what was the question?” (EMNLP 2020)
  • Greedy decoding commits to the best next token and can miss the best sequence; beam search keeps the beam_size best prefixes and finds it in the worked example, 0.36 against greedy’s 0.35 (test_hand_example_beam_beats_greedy, test_hand_example_greedy_commits_to_a).
  • All k * V extensions are ranked together, so the beam never holds more than beam_size rows (test_beam_never_holds_more_than_beam_size).
  • A beam of 1 is greedy decoding, and a beam of at least VLV^L is exact search; both are checked against independent oracles (test_beam_one_is_greedy, test_exhaustive_beam_equals_brute_force).
  • The model is only a step function: its state is gathered by parent with select_state, so the same search decodes a toy table and a real Seq2Seq (test_state_follows_parents_in_a_real_search, test_beam_decodes_a_seq2seq).
  • Summed log-probability prefers short outputs; dividing by len ** length_penalty can change the winner (test_length_penalty_changes_the_winner).
Terminal window
ol start L4.4 # stubs beam.py; prints your test path and rung (R5)
ol tests L4.4 # the course tests
# write your oracle tests in python/tests/l4-4-beam/ (section 4 lists what to cover), then:
ol check L4.4 # course tests and the mutation grade of your tests
ol mutate L4.4 # the full grade, cached by your test files' hash
ol diff L4.4 # after passing: your code against the reference

Your seq2seq model (L4.1) can translate, but only greedily: at each step it takes the single most likely token and never looks back. On the reversal task that is fine, because the model is nearly certain at every step. On the dates task of MS-L4, where the first digit of the year is often uncertain, a greedy decoder commits to a likely-looking first token and then has to live with a worse whole output. This module writes the search that fixes it, once, as a function of a step function, so that later the zoo (L6.7) uses the same code for seq2seq and transformer checkpoints. The model and the search are separate pieces with a narrow interface between them, which is how production decoders (vLLM’s beam search, Hugging Face’s generate) are built too.

SymbolMeaningType / shape
VVvocabulary sizeint
y1:ty_{1:t}a prefix: the tokens generated after boslist of ids
p(v∣y1:t)p(v \mid y_{1:t})the model’s next-token distribution, softmax of its logitsfloat64[V]
ℓ(y1:t)=∑i≤tlog⁡p(yi∣y1:i−1)\ell(y_{1:t}) = \sum_{i \le t} \log p(y_i \mid y_{1:i-1})the prefix’s log-probability (logprob)float
KKbeam_sizeint
kknumber of live hypotheses at this step, k≤Kk \le Kint
LLmax_len, the most tokens generatedint
α\alphalength_penaltyfloat ≥0\ge 0
s(y)=ℓ(y)/∣y∣αs(y) = \ell(y) / \lvert y \rvert^\alphathe score a finished hypothesis is ranked byfloat

A decoder model defines P(y)=∏ip(yi∣y1:i−1)P(y) = \prod_i p(y_i \mid y_{1:i-1}) over sequences that end with eos. The best output is arg⁡max⁡yℓ(y)\arg\max_y \ell(y). There are VLV^L sequences of length LL; with V=32 000V = 32\,000 and L=20L = 20 exhaustive search is impossible, and dynamic programming does not help because the model’s state depends on the whole prefix. Every practical decoder is an approximation.

Greedy decoding takes yt=arg⁡max⁡vp(v∣y1:t−1)y_t = \arg\max_v p(v \mid y_{1:t-1}) at every step (ties to the lowest id) until eos or LL tokens. It is a beam of size 1. Its failure is commitment: a token that is slightly more likely now can lead to a much less likely continuation. The worked example in section 3 is exactly that case.

Beam search keeps up to KK prefixes. One step, in the contract’s op order:

  1. Call step_fn(state, y_prev); it returns logits [k, V] for the kk live hypotheses and the new state.
  2. Normalize each row with log_softmax (M09.2): logits are scores, not log-probabilities.
  3. Every extension gets ℓ(y1:t)+log⁡p(v∣y1:t)\ell(y_{1:t}) + \log p(v \mid y_{1:t}): a [k, V] table of candidates.
  4. Rank all kVk V candidates together, highest first, ties to the lower flat index bV+vb V + v (lower beam, then lower token id), and keep the first K−(finished so far)K - (\text{finished so far}) finite ones.
  5. A kept candidate ending in eos is finished and leaves the beam, keeping its slot: the beam narrows by one. The rest are the next live hypotheses.
  6. The state rows of the new hypotheses are their parents’ rows: select_state(new_state, parents).

The search stops when no hypothesis is live, or after LL tokens, when the live ones are returned unfinished. The result is the finished hypotheses ranked by score.

select_state is what makes the search generic. A state is any nesting of tuples, NamedTuples, lists, and dicts whose array leaves (anything with a shape of at least one dimension, numpy arrays and Tensors alike) have the hypothesis on axis 0. It gathers leaf[parents] for every leaf, rows may repeat (two children of one parent), and everything else is passed through. A Python list is a container, not an array: its elements are searched for arrays, and its own order is not gathered.

Two equivalences pin the algorithm down. With K=1K = 1 it is greedy. With K≥VLK \ge V^L nothing is ever pruned: at step tt there are at most VtV^{t} candidates, all kept, so the result is every possible output in score order, which a test enumerates by brute force.

Every token multiplies the probability by a number below 1, so ℓ\ell always prefers shorter outputs. Wu et al. (GNMT) rank finished hypotheses by ℓ(y)/∣y∣α\ell(y) / \lvert y \rvert^\alpha; here ∣y∣\lvert y \rvert counts eos. α=0\alpha = 0 is the raw log-probability and α=1\alpha = 1 the mean per token. During the search all live hypotheses have the same length, so ranking them by ℓ\ell or by ss is the same thing; the penalty only matters when finished hypotheses of different lengths are compared.

Three tokens, eos = 0, a = 1, b = 2, and a model that is a table of next-token probabilities:

prefixp(eos)p(\text{eos})p(a)p(a)p(b)p(b)
(empty)0.10.50.4
a0.20.10.7
b0.90.050.05

Greedy (L=2L = 2): a (0.5), then after a the best is b (0.7). Output a b, unfinished, P=0.35P = 0.35, ℓ=ln⁡0.35=−1.0498\ell = \ln 0.35 = -1.0498.

Beam (K=2K = 2, L=2L = 2, α=0\alpha = 0). Step 1 ranks a 0.5, b 0.4, eos 0.1 and keeps a and b. Step 2 ranks all six extensions together:

candidateprobability
b eos0.4×0.9=0.360.4 \times 0.9 = 0.36
a b0.5×0.7=0.350.5 \times 0.7 = 0.35
a eos0.5×0.2=0.100.5 \times 0.2 = 0.10
a a0.05
b a, b b0.02 each

It keeps b eos (finished) and a b, which is cut at L=2L = 2 and returned unfinished. Result: b eos with ℓ=ln⁡0.36=−1.0217\ell = \ln 0.36 = -1.0217, then a b with ln⁡0.35\ln 0.35. The model was called twice, with 1 and then 2 rows. With α=1\alpha = 1 the scores are ln⁡0.36/2=−0.5108\ln 0.36 / 2 = -0.5108 and ln⁡0.35/2=−0.5249\ln 0.35 / 2 = -0.5249, the same order. These are test_hand_example_beam_beats_greedy, test_hand_example_greedy_commits_to_a, and test_hand_example_length_penalty.

@dataclass
class Hypothesis:
tokens: list[int]; logprob: float; score: float; finished: bool
def select_state(state: Any, idx: ArrayLike) -> Any: ...
def beam_search(step_fn, init_state, bos: int, eos: int, beam_size: int, max_len: int,
length_penalty: float = 1.0) -> list[Hypothesis]: ...
def greedy_decode(step_fn, init_state, bos: int, eos: int, max_len: int) -> Hypothesis: ...
# step_fn(state, y_prev int64 [k]) -> (logits [k, V], new_state)
TestKINDChecksWhy it matters downstream
test_hand_example_beam_beats_greedyunitsection 3: b eos (0.36) then a b (0.35), calls [1, 2]you and the test agree on the algorithm
test_hand_example_greedy_commits_to_aunitgreedy gives a b, score equal to logprobthe baseline MS-L4 compares against
test_hand_example_length_penaltyunitscores ln⁡0.36/2\ln 0.36 / 2 and ln⁡0.35/2\ln 0.35 / 2the penalty counts eos
test_beam_one_is_greedydifferentiala beam of 1 equals the test’s own greedy loop on 6 prefix-dependent modelsgreedy is a special case, not a second code path
test_exhaustive_beam_equals_brute_forcedifferentialall 31 outputs of a V=3V = 3, L=4L = 4 model, by score, α∈{0,1}\alpha \in \{0, 1\}the search is exact when nothing is pruned
test_small_beam_is_a_prefix_of_the_exhaustive_ranking_at_step_onepropertywith L=1L = 1 the beam is the KK most likely first tokensranking over all candidates
test_beam_never_holds_more_than_beam_sizepropertythe step function never sees more than KK rowsthe cost bound of beam search
test_eos_finishes_and_shrinks_the_beampropertyeos once and last; logprob equals the path sum; the beam never grows backfinished outputs are final
test_search_stops_when_the_last_hypothesis_finishesboundaryexactly two calls when eos is certain at step 2no call with zero rows
test_max_len_returns_unfinishedboundarya model that never says eos gives KK unfinished outputs of length LLdecoding always terminates
test_length_penalty_changes_the_winnerunitα=0\alpha = 0 picks a eos, α=1\alpha = 1 picks b b b eos; results ranked by scorewhy translation systems use a penalty
test_ties_go_to_the_lowest_idboundaryuniform logits keep tokens 0 and 1 firstthe same order in Python and Rust
test_logits_need_not_be_normalizedpropertyadding a constant per row changes nothingmodels return logits
test_masked_tokens_are_never_chosenboundary-inf is never kept, even if the beam is not fullgrammar and stop-list masks
test_logprob_is_the_sum_along_the_pathpropertylogprob and score recomputed from the modelrescoring and comparison across searches
test_select_state_keeps_the_structureunitNamedTuple, dict, list, repeated rows, non-array leaves; a set raisesany model’s state works
test_state_follows_parents_in_a_real_searchpropertyevery state row extends a prefix the search keptthe step function scores the right prefixes
test_validationboundarybad sizes, wrong logits shape, nothing finitecaller bugs fail loudly
test_beam_decodes_a_seq2seqdifferentialover L4.1’s Seq2Seq with Luong attention: beam 1 equals greedy, every hypothesis’s logprob equals teacher-forced rescoringthe zoo’s call site

Rung R5 asks for oracles: tests whose expected values come from an independent computation, not from your implementation. For beam search the oracle is brute force: write a tiny model (a table, or logits from a hash of the prefix), enumerate every output for V=3V = 3 and L=4L = 4, score each with your own log-softmax, and compare with an exhaustive beam. Add greedy equivalence, the section 3 numbers, a model that never says eos, a mask, ties, and select_state on a NamedTuple inside a dict. Keep the state an array (an object array of prefixes works): a Python list is searched for arrays, not gathered. Import only tinyllm.infer.beam. ol check L4.4 requires a mutation score of at least 0.80 with every pitfall fault killed.

PitfallSymptomCaught by
1. keeping the best KK tokens of each hypothesisthe beam grows to K2K^2 rows, then moretest_beam_never_holds_more_than_beam_size (mutant s01)
2. adding raw logits instead of log-probabilitiesrows with a larger logit scale win; disagrees with brute forcetest_logits_need_not_be_normalized, test_exhaustive_beam_equals_brute_force (mutant s02)
3. an eos hypothesis that keeps goingoutputs with tokens after eos, a search that never ends earlytest_eos_finishes_and_shrinks_the_beam (mutant s04)
4. dropping the live hypotheses at max_lenan empty result for a model that rarely says eostest_max_len_returns_unfinished (mutant s05)
5. keeping -inf candidates to fill the beammasked tokens appear in the outputtest_masked_tokens_are_never_chosen (mutant s09)
6. state rows not gathered by parentthe model scores one prefix and the search records anothertest_state_follows_parents_in_a_real_search, test_beam_decodes_a_seq2seq (mutant s03)
eos left out of the lengthscores disagree with the formulatest_hand_example_length_penalty (mutant s06)
results ranked by logprob, not scorethe length penalty has no effecttest_length_penalty_changes_the_winner (mutant s07)
ties to the highest indexPython and Rust disagree on near-uniform rowstest_ties_go_to_the_lowest_id (mutant s08)
survivors finalized one step earlyoutputs one token shorttest_max_len_returns_unfinished (mutant s11)
a NamedTuple rebuilt as a tuplethe decoder’s state.h access failstest_select_state_keeps_the_structure (mutant s13)
lists not searchedan LSTM’s per-layer state list is never reorderedtest_select_state_keeps_the_structure (mutant s14)
greedy scored with length penalty 1greedy_decode’s score is not its logprobtest_hand_example_greedy_commits_to_a (mutant s15)
the model called after every hypothesis finisheda call with zero rows, wasted worktest_search_stops_when_the_last_hypothesis_finishes (mutant s16)

| Forward | L5.5 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackM09.2log_softmax normalizes each step’s logits
BackL4.1Seq2Seq.decode_step is the step function of the course test, its DecoderState the state
BackL4.3the Luong decoder carries feed, one more Tensor for select_state to gather
ForwardL6.7the zoo decodes seq2seq and transformer checkpoints with beam search and reports beam against greedy
ForwardL8.1the sampler is the other way to pick tokens; it shares the tie rule

If you skip this module, L6.7 stops with BLOCKED ... needs L4.4 once it lands: build it, or pass --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
beam_searchHugging Face generate(num_beams=...)batched beams across many inputs, early_stopping modes, num_return_sequencestransformers/generation/utils.py (_beam_search), beam_search.py (BeamHypotheses)
select_statethe KV-cache reordergathers the attention cache by beam index every step; with paged attention only block tables are copied_reorder_cache in HF models; vLLM’s beam search over forked sequences
length_penaltyGNMT’s ((5+∣y∣)/6)α((5 + \lvert y \rvert)/6)^\alpha and coverage penaltya smoother penalty and a term for source words never attendedWu et al. 2016, section 7; fairseq sequence_generator.py
exact search for tiny VLV^Lexact decoding studiesshows the most likely output is often empty: the beam’s errors helpStahlberg and Byrne, “On NMT Search Errors and Model Errors” (2019)